SpecMem / harness /run_accept.py
inweriok's picture
Initial release: SpecMem harness (code only, credentials-free)
a484e22 verified
Raw
History Blame Contribute Delete
11.8 kB
"""POC acceptance experiment: 3 memory arms across simulated users/sessions.
Pipeline:
1. Build simulated users -> ordered (session-major) instance stream.
2. Generate the genuine greedy target tool call for every unique query from
the served gpt-oss-120b (concurrent; cached by exact query string).
3. Replay the stream through each arm. For each instance an arm first DRAFTS
(from its current memory), we score token-LCP accept vs the target, then
the arm OBSERVES the target (growing its store). static_global observes
only during warmup (session 0) then freezes -- ToolSpec behaviour.
4. Aggregate Mean Accepted Tokens (MAT) and acceptance rate by (arm, session)
and dump results/accept_results.json.
"""
from __future__ import annotations
import argparse
import json
import os
from concurrent.futures import ThreadPoolExecutor
from collections import defaultdict
from pathlib import Path
from . import metrics
from .client import ToolClient
from .data import load_bfcl, load_sealtools, load_tau2
from .memory import Embedder, NoMemory, PersonalMemory, StaticGlobal
from .simulate import build_users
ROOT = Path(__file__).resolve().parent.parent
RESULTS = ROOT / "results"
# Tokenizer for the token-LCP accept metric: HF hub id by default;
# override with a local snapshot path if running offline.
MODEL_PATH = os.environ.get("SPECMEM_TOKENIZER", "openai/gpt-oss-120b")
def generate_targets(client, instances, workers=16):
"""Return {query: canonical_target_str} for every unique query."""
uniq = {}
for ins in instances:
uniq.setdefault(ins.query, ins.functions)
items = list(uniq.items())
def _one(qf):
q, funcs = qf
call = client.generate_call(q, funcs)
if call is None:
return q, None
return q, metrics.canonical_call_str(call["name"], call["arguments"])
targets = {}
with ThreadPoolExecutor(max_workers=workers) as ex:
for i, (q, tgt) in enumerate(ex.map(_one, items)):
targets[q] = tgt
if (i + 1) % 25 == 0:
print(f" targets {i+1}/{len(items)}", flush=True)
return targets
def _replay(instances, targets, embedder, args, per_instance):
"""Replay one seed's stream through the 3 arms; return per-session scores."""
arms = [NoMemory(), StaticGlobal(),
PersonalMemory(capacity=args.capacity, eviction=args.eviction)]
agg = {a.name: defaultdict(list) for a in arms}
cur_session = -1
for ins in instances:
tgt = targets.get(ins.query)
if tgt is None:
continue
if ins.session != cur_session:
cur_session = ins.session
for a in arms: # freeze static after warmup
if isinstance(a, StaticGlobal) and cur_session == 1:
a.freeze()
for a in arms:
draft = a.draft(ins.query, ins.functions, ins.user_id, embedder)
sc = metrics.score(draft, tgt)
agg[a.name][ins.session].append(sc)
if a.name == "personal_memory" and per_instance is not None:
per_instance.append({
"user": ins.user_id, "session": ins.session,
"sig": ins.signature_id, "novel": ins.novel,
"accept": sc["accept_length"], "tlen": sc["target_len"]})
call_name, call_args = _parse_target(tgt)
for a in arms:
if isinstance(a, StaticGlobal):
a.observe(ins.query, ins.functions, ins.user_id,
call_name, call_args, embedder)
elif isinstance(a, PersonalMemory):
a.observe(ins.query, ins.functions, ins.user_id,
call_name, call_args, embedder)
if ins.session == 0:
a.seed_shared(ins.query, call_name, call_args, embedder)
return agg
def run(args):
RESULTS.mkdir(exist_ok=True)
# Tokenizer used for the token-LCP accept metric: use the served model's
# own tokenizer so acceptance reflects what a spec decoder for THAT model
# would see. Defaults to gpt-oss for backward compatibility.
metrics.get_tokenizer(args.model_path or MODEL_PATH)
bench = getattr(args, "benchmark", "bfcl")
if bench == "tau2":
raise SystemExit(
"REJECTED DESIGN: --benchmark tau2 previously extracted decision "
"points from tau2-bench's SHIPPED reference trajectories, which "
"were generated with GPT-4.1 as the agent — an off-policy "
"target-substitution bug. Use harness/tau2_live.py to generate canonical "
"traces with the real served model + a live user simulator "
"(requires OPENAI_API_KEY), then score with its replay mode.")
tasks = {"bfcl": load_bfcl, "sealtools": load_sealtools}[bench]()
client = ToolClient(url=args.url, model=args.model)
if not client.ping():
raise SystemExit(f"served model not reachable at {args.url}")
embedder = Embedder()
seeds = list(range(args.seed, args.seed + args.n_seeds))
arm_names = ["no_memory", "static_global", "personal_memory"]
agg = {a: defaultdict(list) for a in arm_names} # pooled over seeds
per_seed_overall = {a: [] for a in arm_names} # MAT per seed
per_instance = []
n_instances_total = n_unique_total = n_none_total = 0
for si, sd in enumerate(seeds):
instances = build_users(
tasks, n_users=args.users, tasks_per_user=args.tasks_per_user,
n_sessions=args.sessions,
queries_per_session=args.queries_per_session, seed=sd)
instances.sort(key=lambda x: (x.session, x.user_id))
n_instances_total += len(instances)
print(f"[seed {sd}] {len(instances)} instances; generating targets ...",
flush=True)
targets = generate_targets(client, instances, workers=args.workers)
n_none = sum(1 for v in targets.values() if v is None)
n_unique_total += len(targets)
n_none_total += n_none
print(f"[seed {sd}] {len(targets)} unique queries, {n_none} no-call",
flush=True)
seed_agg = _replay(instances, targets, embedder, args,
per_instance if si == 0 else None)
for a in arm_names:
for s, xs in seed_agg[a].items():
agg[a][s].extend(xs)
post = [x for s, xs in seed_agg[a].items() if s > 0 for x in xs]
if post:
per_seed_overall[a].append(
sum(x["accept_length"] for x in post) / len(post))
summary = _summarize(agg, args)
overall = _overall(agg, warmup_session=0) # post-warmup aggregate
for a in arm_names: # add cross-seed std of MAT
vals = per_seed_overall[a]
if vals and a in overall:
mean = sum(vals) / len(vals)
var = sum((v - mean) ** 2 for v in vals) / len(vals)
overall[a]["MAT_seed_std"] = round(var ** 0.5, 3)
overall[a]["n_seeds"] = len(vals)
out = {
"config": vars(args),
"seeds": seeds,
"n_instances": n_instances_total,
"n_unique_queries": n_unique_total,
"n_no_toolcall": n_none_total,
"summary": summary,
"overall_post_warmup": overall,
}
tag = args.tag + "_" if args.tag else ""
(RESULTS / f"{tag}accept_results.json").write_text(json.dumps(out, indent=2))
(RESULTS / f"{tag}personal_per_instance.json").write_text(
json.dumps(per_instance, indent=2))
_write_csv(summary, args.sessions, tag)
print("\n=== Mean Accepted Tokens (MAT) by session ===", flush=True)
_print_table(summary, args.sessions)
print("\n=== Overall (sessions >= 1) ===", flush=True)
for arm, v in overall.items():
print(f" {arm:>16}: MAT={v['MAT']:.2f} "
f"accepted_frac={v['accepted_frac']:.3f} "
f"exact_rate={v['exact_rate']:.3f} n={v['n']}", flush=True)
print("\nWrote results/accept_results.json + accept_by_session.csv", flush=True)
def _overall(agg, warmup_session=0):
out = {}
for arm, per_sess in agg.items():
scores = [x for s, xs in per_sess.items() if s > warmup_session
for x in xs]
if not scores:
continue
n = len(scores)
out[arm] = {
"n": n,
"MAT": round(sum(x["accept_length"] for x in scores) / n, 3),
"accepted_frac": round(sum(x["accepted_frac"] for x in scores) / n, 4),
"exact_rate": round(sum(1 for x in scores if x["exact"]) / n, 4),
}
return out
def _write_csv(summary, n_sessions, tag=""):
lines = ["arm,session,n,MAT,accepted_frac,exact_rate"]
for arm in summary:
for s in range(n_sessions):
v = summary[arm].get(str(s))
if v:
lines.append(f"{arm},{s},{v['n']},{v['MAT']},"
f"{v['accepted_frac']},{v['exact_rate']}")
(RESULTS / f"{tag}accept_by_session.csv").write_text("\n".join(lines) + "\n")
def _parse_target(tgt: str):
d = json.loads(tgt)
return d["name"], d.get("arguments", {})
def _summarize(agg, args):
summary = {}
for arm, per_sess in agg.items():
summary[arm] = {}
for s, scores in per_sess.items():
n = len(scores)
mat = sum(x["accept_length"] for x in scores) / n
frac = sum(x["accepted_frac"] for x in scores) / n
exact = sum(1 for x in scores if x["exact"]) / n
summary[arm][str(s)] = {"n": n, "MAT": round(mat, 3),
"accepted_frac": round(frac, 4),
"exact_rate": round(exact, 4)}
return summary
def _print_table(summary, n_sessions):
arms = list(summary.keys())
header = "session | " + " | ".join(f"{a:>16}" for a in arms)
print(header)
print("-" * len(header))
for s in range(n_sessions):
cells = []
for a in arms:
v = summary[a].get(str(s))
cells.append(f"{v['MAT']:>16.2f}" if v else " " * 16)
print(f"{s:>7} | " + " | ".join(cells))
def main():
p = argparse.ArgumentParser()
p.add_argument("--users", type=int, default=6)
p.add_argument("--tasks-per-user", type=int, default=5)
p.add_argument("--sessions", type=int, default=8)
p.add_argument("--queries-per-session", type=int, default=4)
p.add_argument("--capacity", type=int, default=32)
p.add_argument("--eviction", default="lru", choices=["lru", "lfu"])
p.add_argument("--workers", type=int, default=16)
p.add_argument("--seed", type=int, default=0)
p.add_argument("--n-seeds", type=int, default=1)
p.add_argument("--url", default="http://localhost:30000/v1",
help="OpenAI-compatible endpoint of the served model.")
p.add_argument("--model", default="gpt-oss-120b",
help="served-model-name to target.")
p.add_argument("--model-path", default="",
help="local path/HF id for the tokenizer used by the "
"accept metric. Empty -> gpt-oss tokenizer.")
p.add_argument("--tag", default="", help="output filename prefix "
"(e.g. 'phase2' -> results/phase2_accept_results.json). "
"Empty keeps the original POC filenames.")
p.add_argument("--benchmark", default="bfcl",
choices=["bfcl", "sealtools", "tau2"],
help="task pool: BFCL v4 (default), Seal-Tools in-domain "
"test split, or tau2-bench frozen-trajectory decision "
"points.")
run(p.parse_args())
if __name__ == "__main__":
main()