diff --git a/src/post_training/chat_templates/registry.py b/src/post_training/chat_templates/registry.py index c067e2d..03534b3 100644 --- a/src/post_training/chat_templates/registry.py +++ b/src/post_training/chat_templates/registry.py @@ -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*. diff --git a/src/post_training/methods/common.py b/src/post_training/methods/common.py index 7bc50a4..e2ea715 100644 --- a/src/post_training/methods/common.py +++ b/src/post_training/methods/common.py @@ -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 @@ -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] + 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) diff --git a/src/post_training/methods/dpo.py b/src/post_training/methods/dpo.py index 4b2a0ce..a39b07d 100644 --- a/src/post_training/methods/dpo.py +++ b/src/post_training/methods/dpo.py @@ -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, @@ -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 diff --git a/src/post_training/methods/sft.py b/src/post_training/methods/sft.py index 9fe5e59..bd9bb30 100644 --- a/src/post_training/methods/sft.py +++ b/src/post_training/methods/sft.py @@ -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, @@ -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 diff --git a/tests/test_eos_alignment.py b/tests/test_eos_alignment.py new file mode 100644 index 0000000..9f68ee4 --- /dev/null +++ b/tests/test_eos_alignment.py @@ -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, +`` 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(): + """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"