Converter to the Qwen 3.8 chat template (thinking off) with tokenizer round-trip check

Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com>
This commit is contained in:
Kral
2026-10-05 12:03:54 +02:00
parent d5e43e1a83
commit c2d4997566

108
train/to_qwen.py Normal file
View File

@@ -0,0 +1,108 @@
"""Convert accepted trajectories to the Qwen 3.8 chat template (thinking off, Qwen tool call format).
train/.venv/bin/python train/to_qwen.py runs/traj/accepted.jsonl runs/traj/qwen.jsonl
Output per line: id, text (full rendered conversation), assistant_spans ([start, end] character offsets of the
parts the model must learn: after "<think>\\n\\n</think>\\n\\n" up to and with <|im_end|>), n_tokens, messages,
tools, metadata. Checks per record (a failing record is written to <out>.rejected.jsonl with the reason):
- every tool call parses back from the rendered text with the same name and parameter values;
- the text starts like the generation prompt of the same messages (so inference and training match);
- no "</parameter>" or "</tool_call>" inside a parameter value (would break the Qwen format).
"""
import json
import os
import re
import statistics
import sys
from transformers import AutoTokenizer
MODEL_DIR = os.path.expanduser("~/models/Qwen3.8-27B-4bit")
HEAD = "<|im_start|>assistant\n<think>\n\n</think>\n\n"
CALL = re.compile(r"<tool_call>\n<function=([^>\n]+)>\n(.*?)</function>\n</tool_call>", re.S)
PARAM = re.compile(r"<parameter=([^>\n]+)>\n(.*?)\n</parameter>\n", re.S)
def to_template_messages(messages):
out = []
for m in messages:
m = {k: v for k, v in m.items() if k in ("role", "content", "tool_calls", "tool_call_id")}
if m["role"] == "assistant":
m["content"] = m.get("content") or ""
calls = []
for c in m.get("tool_calls") or []:
a = c["function"].get("arguments") or "{}"
a = json.loads(a) if isinstance(a, str) else a
calls.append({"id": c.get("id", ""), "type": "function",
"function": {"name": c["function"]["name"], "arguments": a}})
if calls:
m["tool_calls"] = calls
else:
m.pop("tool_calls", None)
out.append(m)
return out
def value_text(v):
return v if isinstance(v, str) else json.dumps(v)
def check_calls(messages, text):
"""Rendered tool calls == original calls (name and parameter values), in order."""
want = [(c["function"]["name"], {k: value_text(v) for k, v in c["function"]["arguments"].items()})
for m in messages if m["role"] == "assistant" for c in m.get("tool_calls") or []]
body_text = text[text.find("<|im_end|>"):] # skip the system turn: its tool format example looks like a call
got = [(n, {k: v for k, v in PARAM.findall(body)}) for n, body in CALL.findall(body_text)]
return want == got
def spans(text):
out, pos = [], 0
while True:
i = text.find(HEAD, pos)
if i < 0:
return out
start = i + len(HEAD)
end = text.find("<|im_end|>", start) + len("<|im_end|>")
out.append([start, end])
pos = end
def main():
src, dst = sys.argv[1], sys.argv[2]
tok = AutoTokenizer.from_pretrained(MODEL_DIR)
lengths, bad = [], 0
with open(dst, "w") as f, open(dst + ".rejected.jsonl", "w") as rj:
for line in open(src):
rec = json.loads(line)
msgs = to_template_messages(rec["messages"])
why = []
if any(("</parameter>" in value_text(v) or "</tool_call>" in value_text(v))
for m in msgs for c in m.get("tool_calls") or [] for v in c["function"]["arguments"].values()):
why.append("tag_in_parameter")
text = tok.apply_chat_template(msgs, tools=rec["tools"], tokenize=False, enable_thinking=False)
prefix = tok.apply_chat_template(msgs[:2], tools=rec["tools"], tokenize=False,
add_generation_prompt=True, enable_thinking=False)
if not text.startswith(prefix):
why.append("prompt_prefix_differs")
if not check_calls(msgs, text):
why.append("tool_call_roundtrip")
n = len(tok(text, add_special_tokens=False)["input_ids"])
if why:
bad += 1
rj.write(json.dumps({"id": rec["id"], "why": why}) + "\n")
continue
lengths.append(n)
f.write(json.dumps({"id": rec["id"], "task": rec["task"], "teacher": rec["teacher"],
"metadata": rec["metadata"], "repair": rec.get("repair"),
"text": text, "assistant_spans": spans(text), "n_tokens": n,
"messages": msgs, "tools": rec["tools"]}) + "\n")
if lengths:
print(f"{len(lengths)} converted, {bad} rejected; tokens min {min(lengths)} median "
f"{int(statistics.median(lengths))} max {max(lengths)}")
else:
print("nothing converted;", bad, "rejected")
if __name__ == "__main__":
main()