| """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" |
| |
| |
| 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: |
| 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) |
| |
| |
| |
| 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} |
| per_seed_overall = {a: [] for a in arm_names} |
| 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) |
| for a in arm_names: |
| 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() |
|
|