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