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:
108
train/to_qwen.py
Normal file
108
train/to_qwen.py
Normal 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()
|
||||
Reference in New Issue
Block a user