| """Phase 4 — SuffixDecoding across the MAIN TABLE's 4 benchmarks (CPU, cached). |
| |
| Generalizes the tau2-bench SuffixDecoding check (phase4_suffixdecoding.py) to the |
| four standard benchmarks (API-Bank, ToolAlpaca, BFCL, ToolBench) using the SAME |
| multi-session workload construction and the SAME cached greedy targets the main |
| table already uses (results/phase4_mt_targets_{ds}_seed{sd}.json) — so no server |
| and no re-decoding are needed; MAT is directly comparable to phase4_main_table.json. |
| |
| Arms (post-warmup MAT, 3 seeds): |
| static_global frozen (ToolSpec regime) |
| global_evict LIVE global, EMBEDDING-similarity retrieval |
| suffixdecoding LIVE global, TOKEN-SUFFIX-match retrieval (SuffixDecoding) |
| personal_memory ours (LIVE per-user, embedding) |
| |
| Isolates the same question as the tau2 run across four independent benchmarks: |
| does the freshness gain depend on the retrieval mechanism (embedding vs exact |
| token match), or only on the store being live? |
| |
| Run from the repo root: python -m harness.phase4_suffixdecoding_maintable |
| """ |
| from __future__ import annotations |
|
|
| import json |
| import os |
| import statistics as st |
| from collections import defaultdict |
| from pathlib import Path |
|
|
| from . import metrics |
| from .data import load_apibank, load_bfcl, load_toolalpaca, load_toolbench |
| from .memory import (Embedder, GlobalEvict, NoMemory, PersonalMemory, |
| StaticGlobal, SuffixDecodingBaseline, ToolSpecBaseline) |
| from .run_accept import _parse_target |
| 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") |
| DATASETS = {"apibank": load_apibank, "toolalpaca": load_toolalpaca, |
| "bfcl": load_bfcl, "toolbench": load_toolbench} |
| |
| |
| |
| |
| |
| |
| |
| |
| ARMS = ["no_memory", "static_global", "suffixdecoding", "personal_memory"] |
| TASKS_PER_USER = 10 |
| CAP = 48 |
|
|
|
|
| def _make_arms(footprint): |
| return [NoMemory(), StaticGlobal(), |
| SuffixDecodingBaseline(capacity=footprint), |
| PersonalMemory(capacity=CAP, eviction="lru")] |
|
|
|
|
| def _replay(inst, targets, emb, footprint): |
| arms = _make_arms(footprint) |
| agg = {a.name: defaultdict(list) for a in arms} |
| cur = -1 |
| for ins in inst: |
| tgt = targets.get(ins.query) |
| if tgt is None: |
| continue |
| if ins.session != cur: |
| cur = ins.session |
| if cur == 1: |
| for a in arms: |
| if hasattr(a, "freeze"): |
| a.freeze() |
| for a in arms: |
| agg[a.name][ins.session].append(metrics.score( |
| a.draft(ins.query, ins.functions, ins.user_id, emb), tgt)) |
| cn, ca = _parse_target(tgt) |
| for a in arms[1:]: |
| a.observe(ins.query, ins.functions, ins.user_id, cn, ca, emb) |
| if isinstance(a, PersonalMemory) and ins.session == 0: |
| a.seed_shared(ins.query, cn, ca, emb) |
| out = {} |
| for a in arms: |
| post = [x for s, xs in agg[a.name].items() if s > 0 for x in xs] |
| out[a.name] = round(sum(x["accept_length"] for x in post) / |
| max(1, len(post)), 3) |
| return out |
|
|
|
|
| def main(): |
| metrics.get_tokenizer(MODEL_PATH) |
| emb = Embedder() |
| table = {} |
| for ds, loader in DATASETS.items(): |
| tasks = loader() |
| n_users = min(40, len(tasks) // TASKS_PER_USER) |
| footprint = n_users * CAP |
| per_seed = {a: [] for a in ARMS} |
| for sd in (0, 1, 2): |
| cache_f = RESULTS / f"phase4_mt_targets_{ds}_seed{sd}.json" |
| if not cache_f.exists(): |
| print(f" [{ds} seed {sd}] MISSING cache -> skip", flush=True) |
| continue |
| targets = json.loads(cache_f.read_text()) |
| inst = build_users(tasks, n_users=n_users, |
| tasks_per_user=TASKS_PER_USER, n_sessions=12, |
| queries_per_session=6, seed=sd) |
| inst.sort(key=lambda x: (x.session, x.user_id)) |
| res = _replay(inst, targets, emb, footprint) |
| for a in ARMS: |
| per_seed[a].append(res[a]) |
| print(f" [{ds} seed {sd}] " + |
| " ".join(f"{a}={res[a]}" for a in ARMS), flush=True) |
| cells = {a: {"MAT_mean": round(st.mean(per_seed[a]), 3), |
| "MAT_std": round(st.pstdev(per_seed[a]), 3) |
| if len(per_seed[a]) > 1 else 0.0} |
| for a in ARMS if per_seed[a]} |
| sg = cells["static_global"]["MAT_mean"] |
| for a in ARMS: |
| if a in cells and sg: |
| cells[a]["rel_over_static_pct"] = round( |
| 100 * (cells[a]["MAT_mean"] - sg) / sg, 1) |
| table[ds] = {"n_users": n_users, "footprint": footprint, "cells": cells} |
| print(f"=== {ds} done ===", flush=True) |
|
|
| out = {"arms": ARMS, "datasets": list(DATASETS), "table": table, |
| "note": ("Post-warmup MAT (sessions 1-11), 3 seeds, cached greedy " |
| "targets identical to phase4_main_table.json; CPU replay, no " |
| "server. global_evict and suffixdecoding are both LIVE + " |
| "global + size-capped to n_users*48, differing ONLY in " |
| "retrieval (embedding cosine vs exact token-suffix match). " |
| "Real replay outputs; no tuning to a target outcome.")} |
| (RESULTS / "phase4_suffixdecoding_maintable.json").write_text( |
| json.dumps(out, indent=2)) |
| print("\n=== post-warmup MAT (mean over 3 seeds) ===") |
| hdr = "arm".ljust(16) + "".join(d[:9].ljust(11) for d in DATASETS) |
| print(hdr) |
| for a in ARMS: |
| print(a.ljust(16) + "".join( |
| f"{table[d]['cells'][a]['MAT_mean']}".ljust(11) |
| for d in DATASETS if a in table[d]["cells"])) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|