Download code/scripts/teacher.py from darioooooo0o/tiny-agent-112m: direct link, hf CLI and curl.
- Browser
- Download file 12.8 kB
-
https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/teacher.py
- Command line
-
hf download hf://darioooooo0o/tiny-agent-112m/code/scripts/teacher.py
-
curl -L -o teacher.py https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/teacher.py
12.8 kB
| """Teacher trajectories: a strong local model (Qwen3.8-27B via vLLM's OpenAI API) acts as the agent | |
| in our sandbox, with our tools and tool-call syntax. Only episodes that are correct AND grounded | |
| are kept, saved as canonical messages (tiny_agent.chat) with the student's system prompt. | |
| Qwen tends to overthink (thinking can run to tens of thousands of tokens), which would blow both the | |
| time budget and the student's 4k context. By default thinking is switched off in the chat template | |
| (enable_thinking=false) and the model is told to write one or two sentences of reasoning before each | |
| call instead; that text becomes the turn's <think> block in the student data. Calls that hit | |
| max_tokens are counted ("cut"), and the log reports how many episodes survive the same filter the | |
| data build uses (scripts.gen_synth.teacher_reject). | |
| Plain chat in/out, but in Qwen's native tool format: the system prompt carries the same "# Tools" | |
| preamble its chat template renders, Qwen writes <tool_call><function=..><parameter=..> XML (asked | |
| for our JSON instead, it blends the two and ~80% of episodes broke), and the calls are converted | |
| to our canonical JSON calls here. The student's parser (tiny_agent.chat) is unchanged. Tool | |
| results go back as <tool_response> blocks, which is how the template renders tool turns. | |
| $TA_PY scripts/teacher.py --url http://127.0.0.1:8000/v1 --model qwen38 --hours 6 --out $TA_DATA/teacher/run1.jsonl | |
| """ | |
| import argparse | |
| import json | |
| import os | |
| import random | |
| import re | |
| import threading | |
| import time | |
| import urllib.request | |
| from collections import Counter | |
| from concurrent.futures import ThreadPoolExecutor | |
| from scripts.gen_synth import teacher_reject | |
| from tiny_agent.chat import SYSTEM_TA_V1, parse_assistant | |
| from tiny_agent.tasks import TRAIN_KINDS, check, grounded, make_task | |
| from tiny_agent.tools import TOOL_SPECS, Workspace | |
| TEACHER_SEED0 = 5_000_000 # eval 0-99,999 | RL 100,000-999,999 | scripted 1,000,000+ | teacher 5,000,000+ | |
| TEACHER_RULES = """ | |
| Results come back in <tool_response> blocks, in the same order. Paths are relative to /work. | |
| Never guess: every fact in your answer must come from a tool result, and compute results by running | |
| code (e.g. python3 -c), not in your head. Finish by calling submit with only the answer (a number, | |
| name, path, DONE or NOT_FOUND), no extra words.""" | |
| BRIEF_RULE = """ | |
| Before each tool call, write one or two short sentences of plain reasoning: what you will check | |
| and why. Keep it brief.""" | |
| _PARAMS = {"bash": {"command": "string"}, | |
| "read": {"path": "string", "offset": "integer", "limit": "integer"}, | |
| "edit": {"path": "string", "old_string": "string", "new_string": "string", "replace_all": "boolean"}, | |
| "write": {"path": "string", "content": "string"}, | |
| "submit": {"answer": "string"}} | |
| _REQUIRED = {"bash": ["command"], "read": ["path"], "edit": ["path", "old_string", "new_string"], | |
| "write": ["path", "content"], "submit": ["answer"]} | |
| QWEN_TOOLS = [{"type": "function", "function": { | |
| "name": n, "description": d, | |
| "parameters": {"type": "object", "properties": {k: {"type": t} for k, t in _PARAMS[n].items()}, | |
| "required": _REQUIRED[n]}}} for n, _, d in TOOL_SPECS] | |
| # what Qwen3.8's chat template renders when tools are passed (copied from chat_template.jinja) | |
| QWEN_TOOL_PREAMBLE = ( | |
| "# Tools\n\nYou have access to the following functions:\n\n<tools>" | |
| + "".join("\n" + json.dumps(t, ensure_ascii=False) for t in QWEN_TOOLS) + "\n</tools>" | |
| "\n\nIf you choose to call a function ONLY reply in the following format with NO suffix:\n\n<tool_call>\n" | |
| "<function=example_function_name>\n<parameter=example_parameter_1>\nvalue_1\n</parameter>\n" | |
| "<parameter=example_parameter_2>\nThis is the value for the second parameter\nthat can span\nmultiple lines\n" | |
| "</parameter>\n</function>\n</tool_call>\n\n<IMPORTANT>\nReminder:\n- Function calls MUST follow the " | |
| "specified format: an inner <function=...></function> block must be nested within <tool_call></tool_call> " | |
| "XML tags\n- Required parameters MUST be specified\n- You may provide optional reasoning for your function " | |
| "call in natural language BEFORE the function call, but NOT after\n- If there is no function call available, " | |
| "answer the question like normal with your current knowledge and do not tell the user about function calls\n" | |
| "</IMPORTANT>") | |
| _NCALL = re.compile(r"<tool_call>(.*?)(?:</tool_call>|$)", re.S) | |
| _FUNC = re.compile(r"<function=([^>\s]+)>") | |
| _PARAM = re.compile(r"<parameter=([^>\s]+)>\n?(.*?)\n?</parameter>", re.S) | |
| def _coerce(name, key, v): | |
| t = _PARAMS.get(name, {}).get(key, "string") | |
| if t == "integer" and re.fullmatch(r"\s*-?\d+\s*", v): | |
| return int(v) | |
| if t == "boolean" and v.strip().lower() in ("true", "false"): | |
| return v.strip().lower() == "true" | |
| return v | |
| def native_to_json(text): | |
| """Rewrite Qwen's native <function=..><parameter=..> calls as our <tool_call>{json}</tool_call>, | |
| with parameter types from the tool schema. Anything else is left for parse_assistant to judge.""" | |
| def one(m): | |
| body = m.group(1) | |
| f = _FUNC.search(body) | |
| if not f: | |
| return m.group(0) | |
| name = f.group(1) | |
| args = {k: _coerce(name, k, v) for k, v in _PARAM.findall(body)} | |
| return "<tool_call>\n" + json.dumps({"name": name, "arguments": args}, ensure_ascii=False) + "\n</tool_call>" | |
| return _NCALL.sub(one, text.split("<|im_end|>")[0]) | |
| def chat(url, model, messages, max_tokens, temperature, top_p=0.8, top_k=20, thinking=False, timeout=600): | |
| body = json.dumps({"model": model, "messages": messages, "max_tokens": max_tokens, | |
| "temperature": temperature, "top_p": top_p, "top_k": top_k, | |
| "chat_template_kwargs": {"enable_thinking": bool(thinking)}}).encode() | |
| req = urllib.request.Request(f"{url}/chat/completions", data=body, headers={"Content-Type": "application/json"}) | |
| with urllib.request.urlopen(req, timeout=timeout) as r: | |
| out = json.loads(r.read()) | |
| choice = out["choices"][0] | |
| msg = choice["message"] | |
| text = msg.get("content") or "" | |
| reasoning = msg.get("reasoning_content") or msg.get("reasoning") | |
| if reasoning and "<think>" not in text: | |
| text = f"<think>\n{reasoning.strip()}\n</think>\n{text}" | |
| if "</think>" in text and "<think>" not in text: # some templates open <think> in the prompt | |
| text = "<think>\n" + text | |
| return text, out.get("usage", {}), choice.get("finish_reason") == "length" | |
| _OVR = {"mtime": None, "args": None} | |
| def live_args(a): | |
| """Settings can be changed while running (no restart, which would end the teacher window): | |
| write e.g. {"max_tokens": 512, "thinking": 0, "temperature": 0.6} to <out>.overrides.json.""" | |
| path = a.out + ".overrides.json" | |
| try: | |
| mt = os.path.getmtime(path) | |
| except OSError: | |
| return a | |
| if mt != _OVR["mtime"]: | |
| try: | |
| ovr = json.load(open(path)) | |
| b = argparse.Namespace(**vars(a)) | |
| for k, v in ovr.items(): | |
| if k in ("max_tokens", "thinking", "temperature", "top_p", "top_k", "max_turns"): | |
| setattr(b, k, type(getattr(a, k))(v)) | |
| _OVR.update(mtime=mt, args=b) | |
| print("overrides", ovr, flush=True) | |
| except (ValueError, OSError, TypeError) as e: | |
| print("bad overrides file:", e, flush=True) | |
| return a | |
| return _OVR["args"] | |
| def episode(seed, kind, a): | |
| task = make_task(random.Random(seed), kind) | |
| canon = task.messages() | |
| rules = TEACHER_RULES + ("" if a.thinking else BRIEF_RULE) | |
| system = QWEN_TOOL_PREAMBLE + "\n\n" + SYSTEM_TA_V1.split("\nTools:")[0] + rules | |
| conv = [{"role": "system", "content": system}, {"role": "user", "content": task.question}] | |
| toks = cut = n_calls = 0 | |
| with Workspace(task.files) as ws: | |
| for _ in range(a.max_turns): | |
| text, usage, was_cut = chat(a.url, a.model, conv, a.max_tokens, a.temperature, a.top_p, a.top_k, | |
| a.thinking) | |
| toks += usage.get("completion_tokens", 0) | |
| cut += was_cut | |
| n_calls += 1 | |
| msg = parse_assistant(native_to_json(text)) | |
| if not msg["think"] and msg["content"]: | |
| # the brief reasoning written before the calls is the student's <think> block | |
| msg["think"], msg["content"] = msg["content"], "" | |
| conv.append({"role": "assistant", "content": text}) | |
| canon.append(msg) | |
| calls = msg["tool_calls"] | |
| if not calls: | |
| break | |
| results = [] | |
| for c in calls: | |
| results.append(c["error"] if "error" in c else ws.call(c["name"], c["arguments"])) | |
| if ws.submitted is not None: | |
| break | |
| canon.append({"role": "tool", "results": results}) | |
| if ws.submitted is not None: | |
| break | |
| conv.append({"role": "user", "content": "".join(f"<tool_response>\n{r}\n</tool_response>\n" for r in results)}) | |
| ok = check(task, ws.submitted, ws) | |
| return dict(seed=seed, kind=kind, ok=ok, grounded=grounded(task, canon), submitted=ws.submitted, | |
| answer=task.answer, completion_tokens=toks, calls=n_calls, cut=cut, messages=canon) | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--url", default="http://127.0.0.1:8000/v1") | |
| ap.add_argument("--model", default="qwen38") | |
| ap.add_argument("--out", required=True) | |
| ap.add_argument("--n", type=int, default=10**9) | |
| ap.add_argument("--hours", type=float, default=6) | |
| ap.add_argument("--concurrency", type=int, default=48) | |
| ap.add_argument("--max_turns", type=int, default=8) | |
| ap.add_argument("--max_tokens", type=int, default=768, help="per call") | |
| ap.add_argument("--thinking", type=int, default=0, help="1 = let Qwen think (it overthinks)") | |
| ap.add_argument("--temperature", type=float, default=0.7) | |
| ap.add_argument("--top_p", type=float, default=0.8) | |
| ap.add_argument("--top_k", type=int, default=20) | |
| ap.add_argument("--kinds", default=",".join(TRAIN_KINDS)) | |
| ap.add_argument("--seed0", type=int, default=TEACHER_SEED0) | |
| a = ap.parse_args() | |
| os.makedirs(os.path.dirname(a.out), exist_ok=True) | |
| kinds = a.kinds.split(",") | |
| # resume: skip seeds already in the output file | |
| done = set() | |
| if os.path.exists(a.out): | |
| for l in open(a.out): | |
| done.add(json.loads(l)["seed"]) | |
| lock = threading.Lock() | |
| stats, t0 = Counter(), time.time() | |
| deadline = t0 + a.hours * 3600 | |
| f = open(a.out, "a") | |
| def job(i): | |
| if time.time() > deadline: | |
| return | |
| seed = a.seed0 + i | |
| if seed in done: | |
| return | |
| try: | |
| r = episode(seed, kinds[i % len(kinds)], live_args(a)) | |
| except Exception as e: # server hiccup: skip this seed | |
| with lock: | |
| stats["error"] += 1 | |
| return | |
| with lock: | |
| stats["episodes"] += 1 | |
| stats[f"ok_{r['kind']}"] += r["ok"] | |
| stats[f"n_{r['kind']}"] += 1 | |
| why, _ = teacher_reject(r) | |
| stats["usable"] += why is None | |
| stats[f"reject_{why}"] += why is not None | |
| stats["calls"] += r["calls"] | |
| stats["cut_calls"] += r["cut"] | |
| stats["completion_tokens"] += r["completion_tokens"] | |
| f.write(json.dumps(r) + "\n") | |
| f.flush() | |
| if stats["episodes"] % 50 == 0: | |
| el = time.time() - t0 | |
| print(json.dumps({"episodes": stats["episodes"], "usable": stats["usable"], "errors": stats["error"], | |
| "rejects": {k[7:]: v for k, v in stats.items() if k.startswith("reject_")}, | |
| "cut_call_rate": round(stats["cut_calls"] / max(1, stats["calls"]), 3), | |
| "tok_per_call": round(stats["completion_tokens"] / max(1, stats["calls"])), | |
| "tok_s": round(stats["completion_tokens"] / el), | |
| "ok_rate": {k: round(stats[f"ok_{k}"] / max(1, stats[f"n_{k}"]), 2) for k in kinds}, | |
| "elapsed_min": round(el / 60, 1)}), flush=True) | |
| with ThreadPoolExecutor(a.concurrency) as ex: | |
| for i in range(a.n): | |
| if time.time() > deadline: | |
| break | |
| ex.submit(job, i) | |
| # keep the queue short so the deadline is respected | |
| while ex._work_queue.qsize() > a.concurrency * 2: | |
| time.sleep(0.2) | |
| print("final", dict(stats), flush=True) | |
| if __name__ == "__main__": | |
| main() | |