qgallouedec HF Staff commited on
Commit
59617bf
·
verified ·
1 Parent(s): d154aab

Add generation markers to the chat template

Browse files

This PR adds `{% generation %}` / `{% endgeneration %}` markers around the assistant output in `chat_template.jinja`, so `return_assistant_tokens_mask=True` returns the assistant tokens. Without them the mask is all zeros, and assistant-only loss (e.g. TRL SFT with `assistant_only_loss=True`) can't work. Rendering is unchanged.

The `<|im_start|>assistant\n` header is now emitted once before the branches, so it stays outside the mask. The mask covers what the model generates, up to and including `<|im_end|>`:

```
<think>\nr\n</think>\n\nHello<|im_end|>\n
```

Checked: the rendered prompt is identical before and after for 40 combinations (system, multi-turn, tool calls and responses, with and without `tools`, `add_generation_prompt`, `enable_thinking`). The mask covers each assistant turn in multi-turn and tool-call conversations, and the `<|im_end|>` stop token is always inside it.

Files changed (1) hide show
  1. chat_template.jinja +6 -3
chat_template.jinja CHANGED
@@ -40,14 +40,16 @@
40
  {%- set content = content.split('</think>')[-1].lstrip('\n') %}
41
  {%- endif %}
42
  {%- endif %}
 
 
43
  {%- if loop.index0 > ns.last_query_index %}
44
  {%- if loop.last or (not loop.last and reasoning_content) %}
45
- {{- '<|im_start|>' + message.role + '\n<think>\n' + reasoning_content.strip('\n') + '\n</think>\n\n' + content.lstrip('\n') }}
46
  {%- else %}
47
- {{- '<|im_start|>' + message.role + '\n' + content }}
48
  {%- endif %}
49
  {%- else %}
50
- {{- '<|im_start|>' + message.role + '\n' + content }}
51
  {%- endif %}
52
  {%- if message.tool_calls %}
53
  {%- for tool_call in message.tool_calls %}
@@ -69,6 +71,7 @@
69
  {%- endfor %}
70
  {%- endif %}
71
  {{- '<|im_end|>\n' }}
 
72
  {%- elif message.role == "tool" %}
73
  {%- if loop.first or (messages[loop.index0 - 1].role != "tool") %}
74
  {{- '<|im_start|>user' }}
 
40
  {%- set content = content.split('</think>')[-1].lstrip('\n') %}
41
  {%- endif %}
42
  {%- endif %}
43
+ {{- '<|im_start|>' + message.role + '\n' }}
44
+ {%- generation %}
45
  {%- if loop.index0 > ns.last_query_index %}
46
  {%- if loop.last or (not loop.last and reasoning_content) %}
47
+ {{- '<think>\n' + reasoning_content.strip('\n') + '\n</think>\n\n' + content.lstrip('\n') }}
48
  {%- else %}
49
+ {{- content }}
50
  {%- endif %}
51
  {%- else %}
52
+ {{- content }}
53
  {%- endif %}
54
  {%- if message.tool_calls %}
55
  {%- for tool_call in message.tool_calls %}
 
71
  {%- endfor %}
72
  {%- endif %}
73
  {{- '<|im_end|>\n' }}
74
+ {%- endgeneration %}
75
  {%- elif message.role == "tool" %}
76
  {%- if loop.first or (messages[loop.index0 - 1].role != "tool") %}
77
  {{- '<|im_start|>user' }}