Download code/scripts/probe_paraphrase.py from darioooooo0o/tiny-agent-112m: direct link, hf CLI and curl.
- Browser
- Download file 3.59 kB
-
https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/probe_paraphrase.py
- Command line
-
hf download hf://darioooooo0o/tiny-agent-112m/code/scripts/probe_paraphrase.py
-
curl -L -o probe_paraphrase.py https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/probe_paraphrase.py
3.59 kB
| """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 <ckpt> [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 | |