SpecMem / harness /phase4_maintable.py
inweriok's picture
Initial release: SpecMem harness (code only, credentials-free)
a484e22 verified
Raw
History Blame Contribute Delete
12 kB
"""Phase 4 MAIN TABLE (Part A) — 4 arms x 4 real datasets, full metrics.
Rows/arms (reuse existing, no reimplementation):
no_memory = Vanilla AR (speedup reference)
toolspec = faithful ToolSpec (Phase 4.4: confidence-gated retrieval + FSM)
static_global = ToolSpec + static memory (simple frozen proxy)
personal_memory= SpecMem (ours: live, evicting, per-user)
Columns/datasets (all REAL): API-Bank, ToolAlpaca, BFCLv4, ToolBench.
Per (arm, dataset) cell we record:
1. MAT (mean accepted tokens) -- replay
2. tokens/s (real decode throughput) -- wall-clock timing
3. e2e wall-clock speedup vs Vanilla AR -- wall-clock timing (p50)
4. retrieval / write-back overhead ms -- timed embed+lookup / embed+insert
5. memory size (entries) at run end -- replay
6. cold-start vs long-term MAT curve -- per-session MAT (replay)
Targets are the served gpt-oss-120b's greedy calls (real, cached per dataset+seed).
Run from the repo root: python -m harness.phase4_maintable --url http://localhost:30000/v1
"""
from __future__ import annotations
import argparse
import json
import os
import random
import time
from collections import defaultdict
from pathlib import Path
import requests
from . import metrics
from .client import ToolClient
from .data import (load_apibank, load_bfcl, load_toolalpaca, load_toolbench)
from .memory import (Embedder, NoMemory, PersonalMemory, StaticGlobal,
ToolSpecBaseline)
from .run_accept import _parse_target, generate_targets
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")
DATASETS = {"apibank": load_apibank, "toolalpaca": load_toolalpaca,
"bfcl": load_bfcl, "toolbench": load_toolbench}
ARMS = ["no_memory", "toolspec", "static_global", "personal_memory"]
def _make_arms(cap):
return [NoMemory(), ToolSpecBaseline(), StaticGlobal(),
PersonalMemory(capacity=cap, eviction="lru")]
def _mem_size(arm):
if isinstance(arm, PersonalMemory):
return sum(len(s) for s in arm.stores.values())
if isinstance(arm, (StaticGlobal, ToolSpecBaseline)):
return len(arm.entries)
return 0
def _replay(instances, targets, emb, cap):
"""3-arm+baseline replay; return per-arm MAT, per-session MAT, mem size."""
arms = _make_arms(cap)
agg = {a.name: defaultdict(list) for a in arms}
cur = -1
for ins in instances:
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:
v = agg[a.name]
post = [x for s, xs in v.items() if s > 0 for x in xs]
by_sess = {s: round(sum(x["accept_length"] for x in xs) / len(xs), 3)
for s, xs in sorted(v.items()) if xs}
out[a.name] = {
"MAT": round(sum(x["accept_length"] for x in post) /
max(1, len(post)), 3),
"accepted_frac": round(sum(x["accepted_frac"] for x in post) /
max(1, len(post)), 3),
"by_session": by_sess, "mem_entries": _mem_size(a)}
return out
def _timed_gen(chat, prompt, ntok):
t0 = time.perf_counter()
r = requests.post(chat, json={"model": "gpt-oss-120b", "temperature": 0.0,
"max_tokens": ntok, "ignore_eos": True,
"messages": [{"role": "user",
"content": prompt}]}, timeout=180)
dt = (time.perf_counter() - t0) * 1000
r.raise_for_status()
return dt
def _wallclock(instances, targets, emb, cap, chat, sample, seed):
"""Real spec-decode wall-clock per arm on a sample: p50/p95 + speedup +
tokens/s. Rebuild arm accept-lengths at each sampled point via replay."""
arms = _make_arms(cap)
pts, cur = [], -1
for ins in instances:
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()
if ins.session > 0:
row = {"query": ins.query, "T": metrics.accept_length(tgt, tgt)[1]}
for a in arms:
row[a.name] = metrics.accept_length(
a.draft(ins.query, ins.functions, ins.user_id, emb), tgt)[0]
pts.append(row)
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)
rng = random.Random(seed)
rng.shuffle(pts)
pts = [p for p in pts if p["T"] >= 2][:sample]
verify = sorted(_timed_gen(chat, pts[i]["query"][:1500], 1)
for i in range(min(12, len(pts))))[6 // 2 or 0]
lat = defaultdict(list)
toks = defaultdict(list)
for row in pts:
prompt, T = row["query"][:1500], max(1, row["T"])
need = {T} | {max(1, T - row[a]) for a in ARMS}
tc = {n: _timed_gen(chat, prompt, n) for n in need}
lat["baseline"].append(tc[T])
for a in ARMS:
ms = verify + tc[max(1, T - row[a])]
lat[a].append(ms)
toks[a].append(T / (ms / 1000.0)) # effective tokens/s
def p(xs, q):
xs = sorted(xs)
return xs[min(len(xs) - 1, int(q * len(xs)))]
base_p50 = p(lat["baseline"], 0.5)
out = {}
for a in ARMS:
out[a] = {"p50_ms": round(p(lat[a], 0.5), 1),
"p95_ms": round(p(lat[a], 0.95), 1),
"tokens_per_s": round(sum(toks[a]) / len(toks[a]), 1),
"speedup_vs_vanilla": round(
p(lat["no_memory"], 0.5) / p(lat[a], 0.5), 3)}
out["_baseline_p50_ms"] = round(base_p50, 1)
return out
def _overhead(emb, chat):
"""Real retrieval (embed+NN over 48-entry store) and write-back (embed+
insert) latency, ms/query, measured separately."""
from .memory import Entry, _best_match
import numpy as np
store = [Entry(emb.embed(f"seed query {i}"), "x") for i in range(48)]
rlat, wlat = [], []
for i in range(150):
q = f"overhead probe query number {i} with args {i*7}"
t0 = time.perf_counter(); e = emb.embed(q); _best_match(e, store)
rlat.append((time.perf_counter() - t0) * 1000)
t0 = time.perf_counter(); e2 = emb.embed(q + " wb"); store.append(Entry(e2, "y"))
wlat.append((time.perf_counter() - t0) * 1000)
return {"retrieval_ms_mean": round(sum(rlat) / len(rlat), 2),
"writeback_ms_mean": round(sum(wlat) / len(wlat), 2)}
def main():
p = argparse.ArgumentParser()
p.add_argument("--url", default="http://localhost:30000/v1")
p.add_argument("--model", default="gpt-oss-120b")
p.add_argument("--datasets", nargs="+", default=list(DATASETS))
p.add_argument("--seeds", nargs="+", type=int, default=[0, 1, 2])
p.add_argument("--wallclock-sample", type=int, default=80)
p.add_argument("--tasks-per-user", type=int, default=10)
p.add_argument("--capacity", type=int, default=48)
p.add_argument("--workers", type=int, default=16)
args = p.parse_args()
chat = args.url.rstrip("/") + "/chat/completions"
metrics.get_tokenizer(MODEL_PATH)
emb = Embedder()
client = ToolClient(url=args.url, model=args.model)
assert client.ping(), f"served model not reachable at {args.url}"
overhead = _overhead(emb, chat)
table = {}
for ds in args.datasets:
tasks = DATASETS[ds]()
n_users = min(40, len(tasks) // args.tasks_per_user)
cfg = {"n_users": n_users, "tasks_per_user": args.tasks_per_user,
"pool": len(tasks)}
print(f"\n=== {ds}: {len(tasks)} tasks -> {n_users} users x "
f"{args.tasks_per_user} ===", flush=True)
per_seed = defaultdict(dict)
by_session_acc = defaultdict(lambda: defaultdict(list))
for sd in args.seeds:
inst = build_users(tasks, n_users=n_users,
tasks_per_user=args.tasks_per_user, n_sessions=12,
queries_per_session=6, seed=sd)
inst.sort(key=lambda x: (x.session, x.user_id))
cache_f = RESULTS / f"phase4_mt_targets_{ds}_seed{sd}.json"
cache = json.loads(cache_f.read_text()) if cache_f.exists() else {}
fmap = {i.query: i.functions for i in inst}
miss = [type("S", (), {"query": q, "functions": fmap[q]})()
for q in {i.query for i in inst} if q not in cache]
if miss:
print(f" [seed {sd}] generating {len(miss)} targets ...", flush=True)
cache.update(generate_targets(client, miss, workers=args.workers))
cache_f.write_text(json.dumps(cache))
res = _replay(inst, cache, emb, args.capacity)
for a in ARMS:
per_seed[a][sd] = res[a]["MAT"]
for s, m in res[a]["by_session"].items():
by_session_acc[a][s].append(m)
if sd == args.seeds[0]:
mem = {a: res[a]["mem_entries"] for a in ARMS}
wc = _wallclock(inst, cache, emb, args.capacity, chat,
args.wallclock_sample, sd)
print(f" [seed {sd}] MAT: " +
" ".join(f"{a}={res[a]['MAT']}" for a in ARMS), flush=True)
import statistics as st
cell = {}
for a in ARMS:
mats = [per_seed[a][sd] for sd in args.seeds]
cell[a] = {
"MAT_mean": round(st.mean(mats), 3),
"MAT_std": round(st.pstdev(mats), 3),
"tokens_per_s": wc[a]["tokens_per_s"],
"speedup_vs_vanilla": wc[a]["speedup_vs_vanilla"],
"wallclock_p50_ms": wc[a]["p50_ms"],
"wallclock_p95_ms": wc[a]["p95_ms"],
"mem_entries": mem[a],
"by_session_MAT": {s: round(sum(v) / len(v), 3)
for s, v in sorted(by_session_acc[a].items())},
}
table[ds] = {"config": cfg, "baseline_p50_ms": wc["_baseline_p50_ms"],
"cells": cell}
out = {"overhead": overhead, "arms": ARMS,
"datasets": list(args.datasets), "table": table,
"note": ("wall-clock via faithful external spec-decode loop, real "
"sglang timers (ignore_eos), NOT engine-integrated; targets "
"= served gpt-oss-120b greedy; BFCL is v4 (superset of the "
"v2 ToolSpec used); ToolBench from OpenBMB Drive.")}
(RESULTS / "phase4_main_table.json").write_text(json.dumps(out, indent=2))
print("\n=== SPEEDUP vs Vanilla AR (p50) ===")
hdr = "arm".ljust(16) + "".join(d[:9].ljust(11) for d in args.datasets)
print(hdr)
for a in ARMS:
row = a.ljust(16) + "".join(
f"{table[d]['cells'][a]['speedup_vs_vanilla']}x".ljust(11)
for d in args.datasets)
print(row)
if __name__ == "__main__":
main()