File size: 6,540 Bytes
a484e22
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
"""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"
# 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}
# Lean arm set: the embedding-vs-token retrieval contrast (global_evict vs
# suffixdecoding) is settled on tau2-bench (phase4_suffixdecoding.json); here we
# corroborate across the 4 standard benchmarks that a LIVE token-match store
# recovers the freshness gain vs the frozen store, like our embedding store.
#   no_memory      floor          (schema draft)
#   static_global  frozen ref     (== toolspec on these workloads)
#   suffixdecoding LIVE token-match (SuffixDecoding)
#   personal_memory ours          (LIVE per-user embedding)
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           # same total size for both global arms
        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()