Skip to content
Open
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
84 changes: 68 additions & 16 deletions tests/datasets/test_glm52_openai_tokenize_fn.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,8 @@
"""GLM-5.2 OpenAI 对话分词行为测试。

TestGlm52Rendering
test_plain_text_matches_hf_and_golden_labels: 普通对话与 HF 模板及慢速 golden 对齐。
test_all_generation_eos_are_supervised_at_assistant_boundaries: 三类停止 token 均按边界参与训练。
test_plain_text_matches_hf_plus_final_eos_and_golden_labels: 普通对话仅比 HF 推理模板多最终 EOS。
test_multiturn_reasoning_defaults_to_preserved_and_can_be_cleared: 默认保留历史推理且支持显式清除。
test_tools_and_loss_switch_follow_template_masking: 工具对话与 loss 开关生成正确标签。
TestGlm52MessageOptions
Expand Down Expand Up @@ -54,8 +55,49 @@ def _label_flags_for_span(tokenizer, text, labels, substring):


class TestGlm52Rendering:
def test_plain_text_matches_hf_and_golden_labels(self, tokenizer, tokenize_fn):
# 验证普通对话的 token 与标签同时对齐 HF 模板和独立慢速实现。
def test_all_generation_eos_are_supervised_at_assistant_boundaries(self, tokenizer, tokenize_fn):
# user/observation 作为轮间停止目标,只有无后继角色的 assistant 才补 endoftext。
messages = [
{"role": "user", "content": "First question"},
{"role": "assistant", "content": "First answer"},
{"role": "user", "content": "Call a tool"},
{
"role": "assistant",
"content": "Calling now",
"tool_calls": [{"function": {"name": "lookup", "arguments": {"key": "value"}}}],
},
{"role": "tool", "content": "result"},
{"role": "assistant", "content": "Unsupervised answer", "loss": False},
{"role": "user", "content": "Final question"},
{"role": "assistant", "content": "Final answer"},
]

tokenized = tokenize_fn({"messages": messages})
rendered = tokenizer.decode(tokenized["input_ids"], skip_special_tokens=False)
hf_rendered = _render_from_hf(tokenizer, messages, add_generation_prompt=False)
stop_ids = {
token: tokenizer.convert_tokens_to_ids(token) for token in ("<|endoftext|>", "<|user|>", "<|observation|>")
}
stop_labels = {
token: [
tokenized["labels"][index]
for index, token_id in enumerate(tokenized["input_ids"])
if token_id == stop_id
]
for token, stop_id in stop_ids.items()
}

assert tokenizer.eos_token == "<|endoftext|>"
assert rendered == hf_rendered + tokenizer.eos_token
assert "First answer<|user|>" in rendered
assert "</tool_call><|observation|>" in rendered
assert rendered.endswith("Final answer<|endoftext|>")
assert stop_labels["<|user|>"] == [-100, stop_ids["<|user|>"], -100]
assert stop_labels["<|observation|>"] == [stop_ids["<|observation|>"]]
assert stop_labels["<|endoftext|>"] == [stop_ids["<|endoftext|>"]]

def test_plain_text_matches_hf_plus_final_eos_and_golden_labels(self, tokenizer, tokenize_fn):
# HF 推理模板不带最终 EOS;XTuner 在最后一个 assistant 末尾显式补齐。
messages = [
{"role": "user", "content": "Hello!"},
{"role": "assistant", "content": "Hi there."},
Expand All @@ -64,16 +106,17 @@ def test_plain_text_matches_hf_and_golden_labels(self, tokenizer, tokenize_fn):
tokenized = tokenize_fn({"messages": messages})
slow_input_ids, slow_labels = glm52_tokenize_fn_slowspeed(tokenizer, messages)

rendered = _render_from_hf(tokenizer, messages, add_generation_prompt=False)
assert tokenized["input_ids"] == tokenizer.encode(rendered, add_special_tokens=False)
hf_rendered = _render_from_hf(tokenizer, messages, add_generation_prompt=False)
rendered = tokenizer.decode(tokenized["input_ids"], skip_special_tokens=False)
assert rendered == hf_rendered + tokenizer.eos_token
assert tokenized["input_ids"] == slow_input_ids
assert tokenized["labels"] == slow_labels
assert (
tokenizer.decode(
[label for label in tokenized["labels"] if label != -100],
skip_special_tokens=False,
)
== "</think>Hi there."
== "</think>Hi there.<|endoftext|>"
)

def test_multiturn_reasoning_defaults_to_preserved_and_can_be_cleared(self, tokenizer, tokenize_fn):
Expand All @@ -86,11 +129,13 @@ def test_multiturn_reasoning_defaults_to_preserved_and_can_be_cleared(self, toke
]

tokenized = tokenize_fn({"messages": messages})
rendered = _render_from_hf(tokenizer, messages, add_generation_prompt=False)
hf_rendered = _render_from_hf(tokenizer, messages, add_generation_prompt=False)
rendered = tokenizer.decode(tokenized["input_ids"], skip_special_tokens=False)
slow_input_ids, slow_labels = glm52_tokenize_fn_slowspeed(tokenizer, messages)

assert "old trace" in rendered
assert tokenized["input_ids"] == tokenizer.encode(rendered, add_special_tokens=False)
assert rendered == hf_rendered + tokenizer.eos_token
assert rendered.count(tokenizer.eos_token) == 1
assert tokenized["input_ids"] == slow_input_ids
assert tokenized["labels"] == slow_labels
assert all(_label_flags_for_span(tokenizer, rendered, tokenized["labels"], "old trace</think>"))
Expand All @@ -99,20 +144,22 @@ def test_multiturn_reasoning_defaults_to_preserved_and_can_be_cleared(self, toke
assert all(_label_flags_for_span(tokenizer, rendered, tokenized["labels"], "Final answer."))

cleared = Glm52ChatMessages(messages=messages).tokenize(tokenizer, clear_thinking=True)
cleared_rendered = _render_from_hf(
cleared_hf_rendered = _render_from_hf(
tokenizer,
messages,
add_generation_prompt=False,
clear_thinking=True,
)
cleared_rendered = tokenizer.decode(cleared["input_ids"], skip_special_tokens=False)
cleared_slow_ids, cleared_slow_labels = glm52_tokenize_fn_slowspeed(
tokenizer,
messages,
clear_thinking=True,
)

assert "old trace" not in cleared_rendered
assert cleared["input_ids"] == tokenizer.encode(cleared_rendered, add_special_tokens=False)
assert cleared_rendered == cleared_hf_rendered + tokenizer.eos_token
assert cleared_rendered.count(tokenizer.eos_token) == 1
assert cleared["input_ids"] == cleared_slow_ids
assert cleared["labels"] == cleared_slow_labels
assert not any(_label_flags_for_span(tokenizer, cleared_rendered, cleared["labels"], "Old answer."))
Expand Down Expand Up @@ -152,17 +199,22 @@ def test_tools_and_loss_switch_follow_template_masking(self, tokenizer, tokenize
]

tokenized = tokenize_fn({"messages": messages, "tools": tools})
rendered = _render_from_hf(
hf_rendered = _render_from_hf(
tokenizer,
messages,
tools=tools,
add_generation_prompt=False,
)
rendered = tokenizer.decode(tokenized["input_ids"], skip_special_tokens=False)
slow_input_ids, slow_labels = glm52_tokenize_fn_slowspeed(tokenizer, messages, tools=tools)

assert tokenized["input_ids"] == tokenizer.encode(rendered, add_special_tokens=False)
assert rendered == hf_rendered + tokenizer.eos_token
assert rendered.count(tokenizer.eos_token) == 1
assert tokenized["input_ids"] == slow_input_ids
assert tokenized["labels"] == slow_labels
observation_id = tokenizer.convert_tokens_to_ids("<|observation|>")
assert tokenized["labels"][tokenized["input_ids"].index(observation_id)] == observation_id
assert tokenized["labels"][-1] == -100
assert not any(
_label_flags_for_span(tokenizer, rendered, tokenized["labels"], '"description": "Gets the weather."')
)
Expand Down Expand Up @@ -214,15 +266,15 @@ def test_default_system_is_inserted_or_replaced(self, tokenizer):
{"role": "system", "content": "Default system instruction."},
*inserted_messages,
]
rendered = _render_from_hf(
hf_rendered = _render_from_hf(
tokenizer,
expected_messages,
add_generation_prompt=False,
)
rendered = tokenizer.decode(inserted["input_ids"], skip_special_tokens=False)

expected_ids = tokenizer.encode(rendered, add_special_tokens=False)
assert inserted["input_ids"] == expected_ids
assert replaced["input_ids"] == expected_ids
assert rendered == hf_rendered + tokenizer.eos_token
assert replaced["input_ids"] == inserted["input_ids"]
assert not any(
_label_flags_for_span(
tokenizer,
Expand Down
60 changes: 42 additions & 18 deletions xtuner/v1/data_proto/messages/glm52_chat.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,8 @@


_MEDIA_TYPES = {"image", "image_url", "video", "video_url", "audio", "audio_url", "input_audio"}
_END_OF_TEXT = "<|endoftext|>"
_NEXT_ROLE_STOP_TOKENS = {"user": "<|user|>", "tool": "<|observation|>"}


def _visible_text(content: Any) -> str:
Expand Down Expand Up @@ -167,10 +169,14 @@ def append(value: str, loss: bool) -> None:
if message.get("role") == "user":
last_user_index = index

# user/observation 角色 token 同时是生成停止目标,其 loss 归属前一个 assistant。
previous_assistant_loss = False
for index, message in enumerate(messages):
role = message.get("role")
if role == "user":
append(f"<|user|>{_visible_text(message.get('content', ''))}", False)
boundary_loss = index > 0 and messages[index - 1].get("role") == "assistant" and previous_assistant_loss
append(_NEXT_ROLE_STOP_TOKENS[role], boundary_loss)
append(_visible_text(message.get("content", "")), False)
elif role == "system":
append(f"<|system|>{_visible_text(message.get('content', ''))}", False)
elif role == "assistant":
Expand All @@ -191,9 +197,17 @@ def append(value: str, loss: bool) -> None:
append(content.strip(), loss)
if message.get("tool_calls"):
append(_render_tool_calls(message["tool_calls"]), loss)
previous_assistant_loss = loss
next_role = messages[index + 1].get("role") if index + 1 < len(messages) else None
if next_role not in _NEXT_ROLE_STOP_TOKENS:
# 没有 user/observation 角色边界时,用标准 EOS 结束 assistant。
append(_END_OF_TEXT, loss)
elif role == "tool":
if index == 0 or messages[index - 1].get("role") != "tool":
append("<|observation|>", False)
boundary_loss = (
index > 0 and messages[index - 1].get("role") == "assistant" and previous_assistant_loss
)
append(_NEXT_ROLE_STOP_TOKENS[role], boundary_loss)
append(_render_tool_result(message.get("content", ""), tools), False)

if add_generation_prompt:
Expand Down Expand Up @@ -255,24 +269,21 @@ def glm52_tokenize_fn_slowspeed(
) -> tuple[list[int], list[int]]:
"""慢速 golden 参考实现:基于 token 级别前缀 diff 对齐 labels。

这份逻辑刻意保持和旧 `golden_tokenize_fn.py` 一致,用来校验 fast path
渲染复用 GLM 多停止边界的 SFT 语义,标签仍通过独立的 token 前缀 diff 计算
1. 先渲染完整对话,得到唯一的 total_ids 参考序列。
2. 对每条需要 loss 的 assistant 消息,渲染其历史前缀并加 generation prompt。
3. 再渲染“历史 + 当前 assistant”,用 token 前缀差得到当前 assistant 应监督的 suffix。
4. 从上次命中位置开始在 total_ids 里顺序查找 suffix,找到后复制到 labels。
"""
# 显式传递确定值,避免 Jinja 将已定义的 None 解释成 False。
hf_kwargs: dict[str, Any] = dict(
tokenize=False,
add_generation_prompt=add_generation_prompt,
# SFT 渲染保持官方轮间格式,并在没有后继角色边界时补标准 EOS。
full_text, _ = render_glm52_chat(
messages,
tools=tools,
add_generation_prompt=add_generation_prompt,
enable_thinking=enable_thinking,
reasoning_effort=reasoning_effort,
clear_thinking=clear_thinking,
)
hf_kwargs.update(kwargs)

full_text = tokenizer.apply_chat_template(messages, **hf_kwargs)
total_ids = tokenizer.encode(full_text, add_special_tokens=False)
labels = [IGNORE_INDEX] * len(total_ids)

Expand All @@ -283,24 +294,37 @@ def glm52_tokenize_fn_slowspeed(
continue

# 历史前缀以 generation prompt 结束;prefix 之后的 token 才是当前 assistant 生成区间。
prompt_kwargs = dict(hf_kwargs)
prompt_kwargs["add_generation_prompt"] = True
prompt_kwargs["tools"] = tools if index == 0 else None
prefix_text = tokenizer.apply_chat_template(messages[:index], **prompt_kwargs)
prefix_text, _ = render_glm52_chat(
messages[:index],
tools=tools if index == 0 else None,
add_generation_prompt=True,
enable_thinking=enable_thinking,
reasoning_effort=reasoning_effort,
clear_thinking=clear_thinking,
)

# 当前截断渲染可能和 full render 不同,例如历史 thinking 会被模板清掉;是否能在 full
# render 中匹配上 suffix,正是 golden 语义的一部分。
message_kwargs = dict(hf_kwargs)
message_kwargs["add_generation_prompt"] = False
message_kwargs["tools"] = tools if index == 0 else None
message_text = tokenizer.apply_chat_template([m.copy() for m in messages[: index + 1]], **message_kwargs)
message_text, _ = render_glm52_chat(
[m.copy() for m in messages[: index + 1]],
tools=tools if index == 0 else None,
add_generation_prompt=False,
enable_thinking=enable_thinking,
reasoning_effort=reasoning_effort,
clear_thinking=clear_thinking,
)

prefix_ids = tokenizer.encode(prefix_text, add_special_tokens=False)
message_ids = tokenizer.encode(message_text, add_special_tokens=False)
content_ids = message_ids[len(prefix_ids) :]
if not content_ids:
continue

next_role = messages[index + 1].get("role") if index + 1 < len(messages) else None
if next_role in _NEXT_ROLE_STOP_TOKENS:
# 截断渲染以 endoftext 结尾;完整对话改用实际的下一角色停止 token。
content_ids[-1] = tokenizer.convert_tokens_to_ids(_NEXT_ROLE_STOP_TOKENS[next_role])

# 在完整 token 序列中做 token 级绝对对齐,避免字符 offset 对特殊 token 的边界解释差异。
for start in range(curr_ptr, len(total_ids) - len(content_ids) + 1):
if total_ids[start : start + len(content_ids)] == content_ids:
Expand Down
Loading