Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 29 additions & 0 deletions src/post_training/chat_templates/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,35 @@ def has_generation_markers(template: str | None) -> bool:
return bool(_GENERATION_OPEN_RE.search(template) and _GENERATION_CLOSE_RE.search(template))


def infer_end_token_from_render(rendered: str, added_tokens: dict[str, int]) -> str | None:
"""The added token a rendered conversation ends on, or ``None``.

Deliberately not named for the eos or for a "stop token": what comes back is
a TEMPLATE-side observation, and under ``qwen3`` it is ``<|im_end|>``, which
is not the model's eos and not a stop token until ``align_generation_eos``
makes it one. Under the ``olmo3-*`` templates it happens to be the eos, but
only because those templates terminate on it.

This is how the turn terminator is established: **by looking at what the
template actually produced**, never from a ``{template: terminator}`` table.
A table goes stale the moment a template changes and the failure is silent —
the model learns to emit one token while the config stops on another.

Pure string logic on purpose, so it is testable without a tokenizer or a
model. The caller renders (see ``build_tokenizer``) and passes the result in
together with ``tokenizer.get_added_vocab()``.

``None`` is the conservative answer and means *change nothing*: it covers a
template whose terminator is not an added token, one that could not be
rendered, and any shape nobody has considered.
"""
tail = (rendered or "").rstrip("\n")
matches = [t for t in added_tokens if t and tail.endswith(t)]
if not matches:
return None
return max(matches, key=len) # longest wins: '<|im_end|>' over 'end|>'


def get_chat_template(name: str) -> str:
"""Return the Jinja source string for the template registered as *name*.

Expand Down
81 changes: 80 additions & 1 deletion src/post_training/methods/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@
from post_training.callbacks.inference_checkpoint import InferenceCheckpointCallback
from post_training.callbacks.mfu import MFUCallback
from post_training.callbacks.throughput import ThroughputCallback
from post_training.chat_templates.registry import get_chat_template
from post_training.chat_templates.registry import get_chat_template, infer_end_token_from_render

if TYPE_CHECKING:
from post_training.config import PostTrainingConfig
Expand Down Expand Up @@ -52,9 +52,88 @@ def build_tokenizer(config: PostTrainingConfig) -> AutoTokenizer:
template_str = get_chat_template(config.data.chat_template)
tokenizer.chat_template = template_str
logger.info("Chat template set to '%s'.", config.data.chat_template)

# The template decides what ends a turn, and SFT trains the model to emit
# exactly that. When it is not the tokenizer's `eos_token`, generation has
# nothing to stop on: the model emits the token it was trained to emit and
# nobody is listening, so it runs to `max_new_tokens` on every prompt.
# `qwen3` ends on `<|im_end|>`; the `olmo3-*` templates already end on
# `eos_token`, for which this is a no-op.
# Set after the pad fallback above, so `pad_token` keeps the model's own eos
# rather than inheriting the turn terminator.
# Rendered through the tokenizer's own `apply_chat_template`, so what we
# inspect is what training will actually produce — and through the public API
# rather than transformers' private jinja helpers.
end_token = None
try:
probe = [{"role": "user", "content": "q"}, {"role": "assistant", "content": "a"}]
rendered = tokenizer.apply_chat_template(probe, tokenize=False)
end_token = infer_end_token_from_render(rendered, tokenizer.get_added_vocab())
except Exception: # noqa: BLE001 - an unrenderable template must not break the run
logger.warning(
"Could not render chat template '%s' to find its turn terminator; "
"leaving eos_token as %s.",
config.data.chat_template,
tokenizer.eos_token,
)
if end_token is not None and end_token != tokenizer.eos_token:
logger.info(
"Chat template '%s' terminates turns with %s, not the tokenizer's "
"eos_token %s. Setting eos_token to %s so generation stops on what "
"the model is trained to emit.",
config.data.chat_template,
end_token,
tokenizer.eos_token,
end_token,
)
tokenizer.eos_token = end_token
return tokenizer


def align_generation_eos(trainer: Any) -> None:
"""Make the trained checkpoint stop on the token the template taught it.

Runs after the trainer exists, because the model's ``generation_config`` is
what ``generate()`` actually reads — **not** ``tokenizer.eos_token``. The two
are independent, and the model's copy comes from the checkpoint: Prelude
ships no ``generation_config.json`` at all and falls back to ``config.json``'s
``eos_token_id``, while OLMo's is an empty ``{}``. Either way the value is the
model's pretraining terminator, which the SFT data never contains.

The terminator is read off the tokenizer, which :func:`build_tokenizer` has
already aligned to the template, and the pretraining eos is read off the
model, which is where it lives. The pretraining eos is KEPT as a secondary
stop id: the model has a prior to emit it that SFT decays without erasing, so
a stray one should stop cleanly rather than render as text. Qwen ships a
two-element list for the same reason.

A no-op when the two already agree — which is the case for every ``olmo3-*``
template, since those terminate on ``eos_token``.
"""
model = getattr(trainer, "model", None)
tokenizer = getattr(trainer, "processing_class", None)
if model is None or tokenizer is None:
return
gc = getattr(model, "generation_config", None)
terminator_id = getattr(tokenizer, "eos_token_id", None)
if gc is None or terminator_id is None:
return

existing = gc.eos_token_id
existing = existing if isinstance(existing, list) else [existing]
existing = [i for i in existing if i is not None]
Comment thread
Neonkraft marked this conversation as resolved.
if existing == [terminator_id]:
return # already correct; touch nothing

gc.eos_token_id = [terminator_id] + [i for i in existing if i != terminator_id]
logger.info(
"generation_config.eos_token_id set to %s (%s from the chat template, "
"then the model's own).",
gc.eos_token_id,
tokenizer.eos_token,
)


def build_model_init_kwargs(config: PostTrainingConfig) -> dict[str, Any]:
"""Return model kwargs forwarded to TRL's model loader."""
dtype = resolve_torch_dtype(config.model.dtype)
Expand Down
2 changes: 2 additions & 0 deletions src/post_training/methods/dpo.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@

from post_training.data.loader import load_and_mix_datasets
from post_training.methods.common import (
align_generation_eos,
build_callbacks,
build_common_training_kwargs,
build_model_init_kwargs,
Expand Down Expand Up @@ -74,5 +75,6 @@ def build_dpo_trainer(config: PostTrainingConfig, run_dir: Path) -> DPOTrainer:
callbacks=build_callbacks(config, run_dir),
)
sanitize_generation_config(trainer)
align_generation_eos(trainer)
prioritize_metric_callbacks(trainer)
return trainer
2 changes: 2 additions & 0 deletions src/post_training/methods/sft.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
from post_training.chat_templates.registry import has_generation_markers
from post_training.data.loader import load_and_mix_datasets
from post_training.methods.common import (
align_generation_eos,
build_callbacks,
build_common_training_kwargs,
build_model_init_kwargs,
Expand Down Expand Up @@ -364,5 +365,6 @@ def build_sft_trainer(config: PostTrainingConfig, run_dir: Path) -> SFTTrainer:
)

sanitize_generation_config(trainer)
align_generation_eos(trainer)
prioritize_metric_callbacks(trainer)
return trainer
212 changes: 212 additions & 0 deletions tests/test_eos_alignment.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,212 @@
"""Tests for aligning the stop token with what the chat template teaches.

SFT trains the model to emit whatever the template writes at the end of a turn.
Under `qwen3` that is `<|im_end|>`; the model's own eos (`<|endoftext|>` on OLMo,
`<eos>` on Prelude) appears nowhere in the data. `generate()` stops on
`model.generation_config.eos_token_id`, NOT on `tokenizer.eos_token`, and that
value comes from the checkpoint — so without alignment the model emits the token
it was trained to emit and nothing is listening, running to `max_new_tokens` on
every prompt.

`test_generate_stops_...` is the one that matters: it drives a real
`model.generate()` on a tiny randomly-initialised model, forcing the terminator
at every step so that stopping is the only variable. No download, CPU, offline.
"""

from __future__ import annotations

import pytest

from post_training.chat_templates.registry import get_chat_template, infer_end_token_from_render
from post_training.methods.common import align_generation_eos

# OLMo's ids, used throughout so the numbers mean something.
TERMINATOR_ID = 100265 # <|im_end|>
NATIVE_EOS_ID = 100257 # <|endoftext|>
ADDED = {"<|im_start|>": 100264, "<|im_end|>": TERMINATOR_ID, "<|endoftext|>": NATIVE_EOS_ID}


class _Tok:
def __init__(self, eos_token="<|im_end|>", eos_token_id=TERMINATOR_ID):
self.eos_token = eos_token
self.eos_token_id = eos_token_id


class _GenConfig:
def __init__(self, eos_token_id):
self.eos_token_id = eos_token_id


class _Model:
def __init__(self, eos_token_id):
self.generation_config = _GenConfig(eos_token_id)


class _Trainer:
def __init__(self, model=None, tokenizer=None):
self.model = model
self.processing_class = tokenizer


# ── deriving the terminator from what the template rendered ────────────


def test_reads_the_terminator_off_the_render():
assert infer_end_token_from_render("<|im_start|>a<|im_end|>", ADDED) == "<|im_end|>"


def test_a_trailing_newline_does_not_hide_it():
"""qwen3 emits '<|im_end|>\\n' after a turn."""
assert infer_end_token_from_render("...<|im_end|>\n", ADDED) == "<|im_end|>"


def test_returns_none_when_the_render_ends_on_ordinary_text():
"""The conservative answer: the caller then changes nothing."""
assert infer_end_token_from_render("...just an answer", ADDED) is None
assert infer_end_token_from_render("", ADDED) is None


def test_longest_match_wins():
added = {"<|im_end|>": 1, "end|>": 2}
assert infer_end_token_from_render("x<|im_end|>", added) == "<|im_end|>"


@pytest.mark.parametrize("name", ["qwen3", "chatml"])
def test_chatml_style_templates_terminate_on_im_end(name):
"""Rendered with the private helper here only because a test may; production
renders through the tokenizer's public `apply_chat_template`."""
utils = pytest.importorskip("transformers.utils.chat_template_utils")
compiled = utils._compile_jinja_template(get_chat_template(name))
rendered, _ = utils._render_with_assistant_indices(
compiled,
[{"role": "user", "content": "q"}, {"role": "assistant", "content": "a"}],
None,
None,
False,
)
assert infer_end_token_from_render(rendered, ADDED) == "<|im_end|>"


@pytest.mark.parametrize("name", ["olmo3", "olmo3-instruct-sft", "olmo3-think-sft", "tulu3"])
def test_olmo_style_templates_terminate_on_the_native_eos(name):
"""These end a final turn on `eos_token`, so alignment must be a NO-OP for
them. Pinned because a change here would silently alter OLMo runs.
"""
utils = pytest.importorskip("transformers.utils.chat_template_utils")
compiled = utils._compile_jinja_template(get_chat_template(name))
rendered, _ = utils._render_with_assistant_indices(
compiled,
[{"role": "user", "content": "q"}, {"role": "assistant", "content": "a"}],
None,
None,
False,
eos_token="<|endoftext|>",
)
assert infer_end_token_from_render(rendered, ADDED) == "<|endoftext|>"


# ── writing it into the checkpoint ────────────────────────────────────


def test_terminator_goes_first_and_the_native_eos_is_kept():
"""The native eos stays as a secondary stop: it was the pretraining
terminator, so the prior to emit it decays under SFT without vanishing."""
trainer = _Trainer(_Model(NATIVE_EOS_ID), _Tok())

align_generation_eos(trainer)

assert trainer.model.generation_config.eos_token_id == [TERMINATOR_ID, NATIVE_EOS_ID]


def test_no_op_when_the_template_already_ends_on_the_models_eos():
"""The olmo3-* case. Left as a scalar, not rewritten to a one-element list,
so an OLMo run is provably untouched."""
trainer = _Trainer(_Model(NATIVE_EOS_ID), _Tok("<|endoftext|>", NATIVE_EOS_ID))

align_generation_eos(trainer)

assert trainer.model.generation_config.eos_token_id == NATIVE_EOS_ID


def test_an_existing_list_is_preserved_without_duplicating():
trainer = _Trainer(_Model([NATIVE_EOS_ID, 999]), _Tok())

align_generation_eos(trainer)

assert trainer.model.generation_config.eos_token_id == [TERMINATOR_ID, NATIVE_EOS_ID, 999]


def test_already_aligned_is_left_alone():
trainer = _Trainer(_Model([TERMINATOR_ID]), _Tok())

align_generation_eos(trainer)

assert trainer.model.generation_config.eos_token_id == [TERMINATOR_ID]


def test_a_none_eos_on_the_model_does_not_produce_a_none_stop_id():
trainer = _Trainer(_Model(None), _Tok())

align_generation_eos(trainer)

assert trainer.model.generation_config.eos_token_id == [TERMINATOR_ID]


@pytest.mark.parametrize(
"trainer", [_Trainer(None, _Tok()), _Trainer(_Model(1), None), _Trainer(None, None)]
)
def test_missing_pieces_are_tolerated(trainer):
align_generation_eos(trainer) # must not raise


# ── the proof: a real generate() call ─────────────────────────────────


def test_generate_stops_on_the_template_terminator_after_alignment():
Comment thread
KonstiNik marked this conversation as resolved.
"""End-to-end on a real model, because everything above only asserts that a
field was set — this asserts that `generate()` obeys it.

A tiny randomly-initialised model is forced to emit the terminator at every
step, so stopping is the only variable. Without alignment it runs the full
`max_new_tokens`; with it, it stops after one token.
"""
torch = pytest.importorskip("torch")
transformers = pytest.importorskip("transformers")

config = transformers.AutoConfig.for_model(
"llama",
vocab_size=100300,
hidden_size=16,
num_hidden_layers=1,
num_attention_heads=2,
intermediate_size=32,
)
model = transformers.AutoModelForCausalLM.from_config(config).eval()

class ForceTerminator(transformers.LogitsProcessor):
def __call__(self, input_ids, scores):
scores[:] = -1e9
scores[:, TERMINATOR_ID] = 0.0
return scores

prompt = torch.tensor([[1, 2, 3]])

def generated_tokens() -> int:
model.generation_config.pad_token_id = NATIVE_EOS_ID
out = model.generate(
prompt,
attention_mask=torch.ones_like(prompt),
max_new_tokens=20,
do_sample=False,
logits_processor=transformers.LogitsProcessorList([ForceTerminator()]),
)
return out.shape[1] - prompt.shape[1]

# before: the model's own eos is the only stop id, and it is never emitted
model.generation_config.eos_token_id = NATIVE_EOS_ID
assert generated_tokens() == 20, "expected it to run to max_new_tokens unaligned"

align_generation_eos(_Trainer(model, _Tok()))

assert model.generation_config.eos_token_id == [TERMINATOR_ID, NATIVE_EOS_ID]
assert generated_tokens() == 1, "expected generation to stop on the terminator"
Loading