feat: stop generation on the token the chat template teaches - #59
Conversation
SFT trains the model to emit whatever the template ends a turn with — <|im_end|> under qwen3 — while generate() reads the model's generation_config, which still holds the pretraining eos that the data never contains. So the model emits the token it was trained to emit and nothing is listening, running to max_new_tokens on every prompt. The terminator is read off the rendered template rather than from a lookup table, so it cannot drift from the thing it describes. The model's own eos is kept as a secondary stop id. A no-op for the olmo3-* templates, which already end a final turn on eos_token.
| return bool(_GENERATION_OPEN_RE.search(template) and _GENERATION_CLOSE_RE.search(template)) | ||
|
|
||
|
|
||
| def terminator_from_render(rendered: str, added_tokens: dict[str, int]) -> str | None: |
There was a problem hiding this comment.
This name was a bit confusing for me. Something like infer_eos_token_from_render would be more intuitive.
There was a problem hiding this comment.
Renamed to infer_end_token_from_render.
I avoided eos_token because what comes back is a template-side observation. Under qwen3 it's <|im_end|>, which is neither the model's eos nor a stop token until align_generation_eos makes it one. Under the olmo3-* templates, it is the eos, but only coincidentally, because those templates terminate on it.
Does the current choice work for you?
There was a problem hiding this comment.
Yes, makes sense :) LGTM, merging.
Neonkraft
left a comment
There was a problem hiding this comment.
Rename the terminator_from_render function and it looks good to merge :)
Review feedback: the old name did not say terminator of what. Avoiding both "eos" and "stop" in the name is deliberate and now stated in the docstring — what comes back is a template-side observation, and under qwen3 it is <|im_end|>, which is neither the model's eos nor a stop token until align_generation_eos makes it one.
Summary
SFT trains the model to emit whatever the chat template ends a turn with —
<|im_end|>underqwen3— butgenerate()stops onmodel.generation_config.eos_token_id, nottokenizer.eos_token, and that value is the model's pretraining eos, which never appears in the SFT data. So the model emits the token it was trained to emit, nothing is listening, and generation runs tomax_new_tokenson every prompt.The fix makes the tokenizer and the model agree with the template.
build_tokenizerrenders a short probe conversation, takes the added token that render ends on as the terminator, and setstokenizer.eos_tokento it;align_generation_eosthen puts that token first ingeneration_config.eos_token_id, keeping the model's original eos behind it as a secondary stop. Reading the terminator off the render, rather than from a hardcoded map of template to terminator, is deliberate: a map goes stale when a template changes, and it fails the same silent way this PR fixes. Templates that already end oneos_token— everyolmo3-*— need no change and get none. Bothsft.pyanddpo.pyuse these helpers, so DPO is covered.Type of change
Validation
pytest tests/— 147 pass, 19 new; ruff 0.9.10 and black 25.1.0 clean.test_generate_stops_on_the_template_terminator_after_alignmentdrives a realmodel.generate()on a tiny randomly-initialised model, forcing the terminator every step so stopping is the only variable: 20 generated tokens before alignment, 1 after. Theolmo3-*no-op is pinned by parametrised tests.