diff --git a/train/to_qwen.py b/train/to_qwen.py new file mode 100644 index 0000000..02512a8 --- /dev/null +++ b/train/to_qwen.py @@ -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 "\\n\\n\\n\\n" up to and with <|im_end|>), n_tokens, messages, +tools, metadata. Checks per record (a failing record is written to .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 "" or "" 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\n\n\n\n" +CALL = re.compile(r"\n\n]+)>\n(.*?)\n", re.S) +PARAM = re.compile(r"\n]+)>\n(.*?)\n\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(("" in value_text(v) or "" 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()