Skip to content

feat: add Qwen3-8B template - #47

Merged
Neonkraft merged 13 commits into
mainfrom
feat/qwen3-sft-template
Aug 25, 2026
Merged

Neonkraft merged 13 commits into
mainfrom
feat/qwen3-sft-template

Conversation

@Neonkraft

Copy link
Copy Markdown
Collaborator

Summary

Adds qwen3-8b.jinja to the chat-template registry directory: Qwen3-8B's upstream template with {% generation %} markers spliced in so it can drive assistant_only_loss=True.

The upstream template concatenates the turn header and the message body into a single emission ('<|im_start|>' + message.role + '\n' + content), so the markers can't simply wrap the assistant branch. The header emission is split out and left outside the markers; {% generation %} opens after it and closes around <|im_end|>, with the trailing \n outside. Result: assistant content, reasoning trace and tool calls are in the loss, prompt tokens and the turn separator are not.

Type of change

  • Bug fix
  • New feature
  • Refactor
  • Performance
  • Documentation
  • Maintenance

Validation

Rendered output compared byte-for-byte against the pristine upstream template across 20 combinations — 5 conversation shapes (single-turn reasoning, multi-turn, no-think, tool call, system+tools) × add_generation_prompt on/off × with/without tools — all identical, confirming the restructuring is loss-mask-only and changes no inference behaviour.

@Neonkraft
Neonkraft requested a review from KonstiNik August 6, 2026 15:32

@KonstiNik KonstiNik left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Nice, thanks for taking care of it!

  • Found one aspect worth noting.
  • Adding it to the tests would also be great I think.
  • Just curious about the naming – are there differences across qwen model sizes?

Comment thread src/post_training/chat_templates/templates/qwen3.jinja
Qwen3 drops <think> from assistant turns at or before the last user query, so history rendered without a reasoning trace was being trained on.
Qwen3 drops <think> from assistant turns at or before the
last user query, so history rendered without a reasoning
trace was being trained on.
Vendors the pristine Qwen/Qwen3-8B chat_template as a fixture and diffs rendered output across the 20-combination matrix.
@Neonkraft

Copy link
Copy Markdown
Collaborator Author

Adding it to the tests would also be great I think.

Done. Also added a test to ensure that messages rendered with the new template are byte-equivalent to those rendered with the upstream template.

Just curious about the naming – are there differences across qwen model sizes?

Fair point. All Qwen3 models except Qwen3-{size}-{Instruct/Thinking}-* seem to have the same template, so I've renamed the template to qwen3.jinja.

@Neonkraft Neonkraft added the enhancement New feature or request label Aug 7, 2026
@Neonkraft

Copy link
Copy Markdown
Collaborator Author

I think we only want to train on the last generation and have the rest as context.

Do we have a strong reason for this? The only explanation I can think of is that the non-final assistant turns wouldn't have the reasoning traces, so the model would be learning to produce a mix of reasoning as well as non-reasoning outputs. On the other hand, wouldn't we be losing a lot of training signal in case of multi-turn conversations?

@KonstiNik

Copy link
Copy Markdown
Member

AFAIK, this is how the Qwen template handles it.
The fix you're referring to would then be done in the data: a sample with multi-turn reasoning would be unrolled into several data samples, stopping at the different turns. So it would create more samples. What are your thoughts on that?

@KonstiNik

Copy link
Copy Markdown
Member

Nice, this looks right to me now. I checked the rendering on the normal shapes and it does what it should.

One thing worth handling before this goes in, and I think it should go into this PR.

Gating the markers makes zero-span rows possible. If a conversation has no assistant turn after the last real user message, there is no generation region and the mask comes out empty. TRL checks for exactly that and raises inside dataset.map:

You're using assistant_only_loss=True, but at least one example has no assistant tokens. This usually means the tokenizer's chat template doesn't generate assistant masks — it may be missing the {% generation %} keyword. […]

So one bad row kills tokenization for the whole mixture, before a single training step, with an error that points at the template instead of the data.

It splits in two, and only half of it is about Qwen:

  • no assistant turn at all ([user], [system, user]) – empty mask under the olmo3-*-sft templates too, I checked, since there is simply nothing to wrap. Pre-existing hole: _sft_row_filter only checks len(messages) > 0, and open-perfectblend has rows like this.
  • an assistant turn, but none after the last real user query – typically ending on a user turn, and also anything with no user turn at all, since ns.last_query_index then falls back to the final index. Qwen-only: trains fine under OLMo, where the dangling user turn is just unsupervised context.

So the solution is part global, part template-aware:

  • require ≥1 assistant turn in _sft_row_filter unconditionally (drop the rows that don't fulfill). That half is a bug fix for OLMo too
  • for the Qwen-only: a capability flag on the template in the registry set only for qwen3, with _sft_row_filter conditioned on it. build_sft_trainer already reads config.data.chat_template just above where it passes the filter, so it's a closure over one boolean
  • log how many rows the filter dropped per dataset. loader.py filters silently today, so a heavily filtered dataset looks the same as a small one – and the weight is applied to the surviving rows, so a silent drop quietly shrinks that dataset's share of the mix. Pre-existing, so happy for it to be a separate PR.

Also worth a zero-span test – the current masking tests cover marker presence and history exclusion, which is why this didn't show up.

One edge case for completeness, though I doubt it shows up in practice: a conversation ending on a tool turn whose assistant caller sits before the last user query ([user, assistant, user, tool]) is also empty. Ordinary tool trajectories are fine. Only worth mentioning because it means the check can't just be "does it end on an assistant turn" – [user, assistant, tool] doesn't either, and that one is fine.

What do you think?

@Neonkraft

Copy link
Copy Markdown
Collaborator Author

a sample with multi-turn reasoning would be unrolled into several data samples, stopping at the different turns.

Yes. This will do.

Gating the markers makes zero-span rows possible

I actually ran into this problem too when I was testing out this template. I also observed the edge cases that you speak of (turn ending in tool use, no assistant turn after the last user turn).

The solutions you propose make sense, but I think the most robust solution would be to simply apply the chat template and filter out rows that have have empty loss masks. Otherwise, we risk an edge case that we haven't considered crashing the run at a later point.

Consider:

ds = load_dataset(DATASET)

# remove_columns keeps the map cache tiny: it holds the flag alone, not a copy of the data.
in_loss = ds.map(
    lambda row: {
        "in_loss": any(
            tok.apply_chat_template(
                row["messages"], return_dict=True, return_assistant_tokens_mask=True
            )["assistant_masks"]
        )
    },
    remove_columns=ds.column_names,
    desc="masking",
)["in_loss"]
in_loss = list(in_loss)

filtered_rows = ds.select([i for i, ok in enumerate(in_loss) if ok])

# Save to disk, point to that in config.yaml

What do you think?

@Neonkraft

Copy link
Copy Markdown
Collaborator Author

See #48 for the full script.

@KonstiNik

Copy link
Copy Markdown
Member

I see the point. Using the template itself to filter is definitely the most robust solution.

I'm just worried that this will drop unconditionally: If a dataset loses 5%, we don't really know what the issue was. Could we pair the mask check with a categorised count of the failure modes we already know: no assistant turn at all, no assistant turn after the last real user query, and an "other" bucket? The mask stays the authority on keep/drop, the categories only say why, and "other" being non-zero is the useful signal, since that's a shape none of us has thought about yet. This way we can catch irregularities.

Moreover, I think we could implement this inside _sft_row_filter rather than as a separate script. build_tokenizer already runs a few lines above where the filter is passed in build_sft_trainer, so it's a closure over the tokenizer. Two things we’d win: the check can't be forgotten, and the filtered set can't go stale — a copy saved to disk is only valid for the template it was built with, so it silently becomes wrong as soon as someone changes data.chat_template.

@Neonkraft

Copy link
Copy Markdown
Collaborator Author

I think we could implement this inside _sft_row_filter rather than as a separate script

Yes, this makes sense. I've changed the implementation of the row filters to accommodate this.

* fix: warn when max_length cuts a supervised span part-way

`any(assistant_masks)` only catches rows whose assistant turn is truncated away
entirely. A row that straddles max_length keeps a non-zero mask, so it stays in
training and teaches an answer cut mid-sentence with no end-of-turn token.

One untruncated render classifies all three cases instead — no assistant tokens,
all of them past the cap, or straddling it — at the same cost. The keep/drop
decision is unchanged.

* feat: add sft.truncated_span_action to drop rows a cut span damages

Warning is right when the cap is wrong, since raising max_seq_length keeps the
data. It is not enough when the sequence length is fixed by memory and dropping
is the only lever left, so make it a switch. Default stays warn.

Drop still warns. data.datasets[].weight multiplies the rows that survive
filtering, so an uneven drop changes a dataset's share of the mixture while the
config still claims the original weights — that has to be in the log rather than
inferred from a row count.

---------

Co-authored-by: Konstantin Nikolaou <knikolaou@icp.uni-stuttgart.de>

@KonstiNik KonstiNik left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM now. Happy to merge it.

@Neonkraft
Neonkraft merged commit c9febad into main Aug 25, 2026
2 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

enhancement New feature or request

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants