"""Wording-robustness probe: the same tasks asked with the template wording (as in the fixed eval) and with fresh phrasings that appear in neither make_task's templates nor vary_question's training bank. A big gap means the policy keys on question templates instead of reading the request. $TA_PY scripts/probe_paraphrase.py [tasks_per_kind=12] [device=xpu] """ import dataclasses, json, random, sys, torch from tokenizers import Tokenizer from tiny_agent.text import DATA from tiny_agent.tasks import make_task from tiny_agent.rollout import make_roller from tiny_agent.checkpoint import load_model def para(t, rng): # fresh wording: in neither make_task's templates nor vary_question's training bank m = t.meta if t.kind == "config_value": return rng.choice([f"Somebody asked me about {m['svc']}'s {m['key']} value. Can you find it?", f"Per its configuration, {m['svc']} has which {m['key']}?"]) if t.kind == "code_constant": return rng.choice([f"The {m['module']} code sets {m['const']} to some number. Which number?", f"Report {m['const']} as set by the {m['module']} source file."]) if t.kind == "csv_lookup": return rng.choice([f"Among the rows of {m['table']}, find {m['name']} and report the {m['col']} column.", f"{m['name']} appears in a CSV called {m['table']}. Its {m['col']}?"]) if t.kind == "long_file": return rng.choice([f"Check the job list for job-{m['job']} and tell me the {m['key']} field.", f"What's recorded as the {m['key']} for the job numbered {m['job']}?"]) if t.kind == "code_search": return rng.choice([f"Some file in the repo has a function named {m['fn']}. Which file?", f"Path of the module that defines {m['fn']}, please."]) if t.kind == "aggregate": return rng.choice([f"Tally the services that run in {m['region']}.", f"Out of all services, how many sit in the {m['region']} region?"]) raise ValueError(t.kind) ck = sys.argv[1] tok = Tokenizer.from_file(f"{DATA}/tokenizer.json") dev = sys.argv[3] if len(sys.argv) > 3 else "xpu" model = load_model(ck, device=dev).eval() kinds = ("config_value", "code_constant", "csv_lookup", "long_file", "code_search", "aggregate") orig, alt = [], [] for k in kinds: for s in range(int(sys.argv[2]) if len(sys.argv) > 2 else 4): t = make_task(random.Random(7_000_000 + s * 31 + len(k)), k) orig.append(t) alt.append(dataclasses.replace(t, question=para(t, random.Random(s)))) roller = make_roller("continuous", model, tok, device=dev, max_len=4096, temperature=1.0, slots=len(orig) * 2) eps = roller.run(orig + alt) n = len(orig) res = {k: [0, 0, 0] for k in kinds} for i, t in enumerate(orig): res[t.kind][0] += eps[i].correct res[t.kind][1] += eps[n + i].correct res[t.kind][2] += 1 for k, (a, b, c) in res.items(): print(f"{k:14s} original {a}/{c} paraphrased {b}/{c}") print("TOTAL original", sum(e.correct for e in eps[:n]), "/", n, " paraphrased", sum(e.correct for e in eps[n:]), "/", n) for i in range(n): if eps[i].correct and not eps[n + i].correct: e = eps[n + i] print("\n--- paraphrase failure:", alt[i].question, "| answer", alt[i].answer, "| submitted", e.ws.submitted) for msg in e.messages[2:]: if msg["role"] == "assistant": print(" THINK:", (msg.get("think") or "")[:120], "| CALLS:", json.dumps(msg["tool_calls"])[:160]) break