File size: 6,454 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
"""Phase 4 — SuffixDecoding baseline validation on tau2-bench (cheap CPU check).

Adds a faithful SuffixDecoding arm (Oliaro et al., NeurIPS 2025; arXiv 2411.04975)
to the acceptance harness and compares it against the arms that isolate the
question it answers:

  no_memory      frozen floor (schema draft)
  toolspec       faithful ToolSpec (frozen, global, confidence-gated + FSM)
  static_global  frozen global proxy
  global_evict   LIVE, global, EMBEDDING-similarity retrieval  <-- contrast A
  suffixdecoding LIVE, global, TOKEN-SUFFIX-match retrieval     <-- contrast A
  personal_memory ours: LIVE, per-user, embedding retrieval     <-- contrast B

Contrast A (global_evict vs suffixdecoding): both live+global+size-capped, they
differ ONLY in retrieval mechanism -> does embedding similarity add value beyond
mere liveness, or is token matching enough?
Contrast B (suffixdecoding/global_evict vs personal_memory): does per-user
partitioning add value on top of a live store?

CPU-only: replays the frozen tau2 decision points against cached on-policy trace
targets (no server). Emits results/phase4_suffixdecoding.json.

Run from the repo root:  python -m harness.phase4_suffixdecoding
"""
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 Task
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")
DOMAINS = ("airline", "retail", "telecom")
ARMS = ["no_memory", "toolspec", "static_global", "global_evict",
        "suffixdecoding", "personal_memory"]


def _make_arms():
    # global arms sized to the same total footprint U*C = 40*48 = 1920.
    return [NoMemory(), ToolSpecBaseline(), StaticGlobal(),
            GlobalEvict(capacity=1920),
            SuffixDecodingBaseline(capacity=1920),
            PersonalMemory(capacity=48, eviction="lru")]


def main():
    metrics.get_tokenizer(MODEL_PATH)
    dp = [json.loads(l) for l in
          (RESULTS / "tau2_live_decision_points.jsonl").read_text().splitlines()]
    tools = {d: json.loads((ROOT / "data" / "tau2" /
                            f"tools_{d}.json").read_text()) for d in DOMAINS}
    tasks = [Task(id=r["id"], query=r["query"], functions=tools[r["domain"]],
                  origin_id=r["id"]) for r in dp]
    targets = {r["query"]: r["target"] for r in dp}
    emb = Embedder()

    per_seed = {a: defaultdict(list) for a in ARMS}   # arm -> session -> [seed MAT]
    post_seed = {a: [] for a in ARMS}                 # arm -> [seed post-warmup MAT]
    for sd in (0, 1, 2):
        inst = build_users(tasks, n_users=40, tasks_per_user=15, n_sessions=12,
                           queries_per_session=6, seed=sd, perturb_prob=0.0)
        inst.sort(key=lambda x: (x.session, x.user_id))
        arms = _make_arms()
        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)
        for a in arms:
            for s, xs in agg[a.name].items():
                per_seed[a.name][s].append(
                    sum(x["accept_length"] for x in xs) / len(xs))
            post = [x for s, xs in agg[a.name].items() if s > 0 for x in xs]
            post_seed[a.name].append(
                sum(x["accept_length"] for x in post) / len(post))
        print(f"seed {sd} done", flush=True)

    curve = {a: {str(s): round(sum(v) / len(v), 3)
                 for s, v in sorted(per_seed[a].items())} for a in ARMS}
    postmat = {a: round(sum(v) / len(v), 3) for a, v in post_seed.items()}
    poststd = {a: round(st.pstdev(v), 3) if len(v) > 1 else 0.0
               for a, v in post_seed.items()}

    def rel(x, base):
        return round(100 * (postmat[x] - postmat[base]) / postmat[base], 1)

    contrasts = {
        "suffixdecoding_over_static_pct": rel("suffixdecoding", "static_global"),
        "global_evict_over_static_pct": rel("global_evict", "static_global"),
        "personal_over_static_pct": rel("personal_memory", "static_global"),
        "embedding_vs_token_gap_pct_of_static": round(
            rel("global_evict", "static_global")
            - rel("suffixdecoding", "static_global"), 1),
        "personal_over_suffixdecoding_pct": rel("personal_memory",
                                                "suffixdecoding"),
        "personal_over_global_evict_pct": rel("personal_memory", "global_evict"),
    }
    out = {"arms": ARMS, "post_warmup_MAT": postmat, "post_warmup_seed_std": poststd,
           "per_session_MAT": curve, "contrasts": contrasts,
           "note": ("tau2-bench 3-domain frozen decision points, cached on-policy "
                    "trace targets, 3 seeds; session 0 = warmup. global_evict and "
                    "suffixdecoding are both LIVE + global + size-capped 1920, "
                    "differing ONLY in retrieval (embedding cosine vs token-suffix "
                    "match) -> isolates retrieval mechanism. All numbers are real "
                    "replay outputs; no tuning to a target outcome.")}
    (RESULTS / "phase4_suffixdecoding.json").write_text(json.dumps(out, indent=2))
    print(json.dumps({"post_warmup_MAT": postmat, "contrasts": contrasts},
                     indent=1))


if __name__ == "__main__":
    main()