File size: 3,589 Bytes
4397e12 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 | """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
|