| 77 |
77 |
return tokens, time.perf_counter() - t0 |
| 78 |
78 |
|
| 79 |
79 |
|
|
80 |
+def judge_outputs(outputs_by_config: dict, prompts: list[dict]) -> dict: |
|
81 |
+ """Quality-level metric: mean per-token logprob of each config's generated |
|
82 |
+ continuation under the bf16 reference model (higher = better). Token-exact |
|
83 |
+ fidelity is incoherent on Metal (1.56%/token prefill/decode flips), so the |
|
84 |
+ judge scores usefulness of the text the system actually produced.""" |
|
85 |
+ import gc |
|
86 |
+ |
|
87 |
+ gc.collect(); mx.clear_cache() |
|
88 |
+ judge, _ = load(HF_MODEL) |
|
89 |
+ scores = {} |
|
90 |
+ for name, outs in outputs_by_config.items(): |
|
91 |
+ vals = [] |
|
92 |
+ for p, toks in zip(prompts, outs): |
|
93 |
+ if len(toks) < 2: |
|
94 |
+ continue |
|
95 |
+ full = p["ids"] + list(toks) |
|
96 |
+ logits = judge(mx.array(full)[None])[0] |
|
97 |
+ sel = logits[len(p["ids"]) - 1 : len(full) - 1].astype(mx.float32) |
|
98 |
+ logprobs = sel - mx.logsumexp(sel, axis=-1, keepdims=True) |
|
99 |
+ idx = mx.array(toks) |
|
100 |
+ tok_lp = mx.take_along_axis(logprobs, idx[:, None], axis=-1) |
|
101 |
+ mx.eval(tok_lp) |
|
102 |
+ vals.append(float(mx.mean(tok_lp).item())) |
|
103 |
+ scores[name] = {"mean_logprob_bf16": sum(vals) / len(vals), "n": len(vals)} |
|
104 |
+ del judge |
|
105 |
+ gc.collect(); mx.clear_cache() |
|
106 |
+ return scores |
|
107 |
+ |
|
108 |
+ |
| 80 |
109 |
def fidelity(a: list[int], b: list[int]) -> float: |
| 81 |
110 |
"""Similarity of two token sequences (difflib ratio — robust to length |
| 82 |
111 |
drift after divergence).""" |
| 134 |
163 |
for tau in ([float(x) for x in args.taus.split(",")] if mode == "margin" else [2.0]): |
| 135 |
164 |
configs.append({"mode": mode, "tau": tau}) |
| 136 |
165 |
|
|
166 |
+ outputs_by_config = {"pure_q4": q4_out, "pure_q8": q8_out} |
| 137 |
167 |
results = [] |
| 138 |
168 |
for cfg in configs: |
| 139 |
169 |
print(f"runtime: mode={cfg['mode']} tau={cfg['tau']} W={args.window} …", flush=True) |
| 140 |
170 |
fid, agg = [], {"tokens": 0, "deferred": 0, "sweeps": 0, "rollbacks": 0, |
| 141 |
171 |
"sweep_s": 0.0, "gen_s": 0.0, "logical_bytes": 0} |
|
172 |
+ cfg_outputs = [] |
| 142 |
173 |
for p, ref in zip(prompts, q8_out): |
| 143 |
174 |
toks, st = generate_deferred( |
| 144 |
175 |
base_model, verify_model, tokenizer, p["ids"], |
| 145 |
176 |
args.max_tokens, cfg["tau"], args.window, cfg["mode"], q8_bytes) |
|
177 |
+ cfg_outputs.append(toks) |
| 146 |
178 |
fid.append(fidelity(toks, ref)) |
| 147 |
179 |
agg["tokens"] += st.tokens_out; agg["deferred"] += st.deferred |
| 148 |
180 |
agg["sweeps"] += st.sweeps; agg["rollbacks"] += st.rollbacks |
| 160 |
192 |
"logical_verify_bytes_per_token": agg["logical_bytes"] / n, |
| 161 |
193 |
"raw": agg, |
| 162 |
194 |
}) |
|
195 |
+ outputs_by_config[f"{cfg['mode']}_tau{cfg['tau']}"] = cfg_outputs |
| 163 |
196 |
r = results[-1] |
| 164 |
197 |
print(f" fidelity={r['fidelity_vs_q8_mean']:.4f} tok/s={r['tokens_per_s']:.1f} " |
| 165 |
198 |
f"defer={r['deferral_rate']:.2f} rollback={r['rollback_rate']:.3f} " |
| 166 |
199 |
f"MB/token(logical)={r['logical_verify_bytes_per_token']/1e6:.0f}", flush=True) |
| 167 |
200 |
|
|
201 |
+ print("judging outputs with bf16 reference …", flush=True) |
|
202 |
+ quality = judge_outputs(outputs_by_config, prompts) |
|
203 |
+ for name, s in quality.items(): |
|
204 |
+ print(f" {name:>18}: mean logprob (bf16 judge) = {s['mean_logprob_bf16']:.4f}", flush=True) |
|
205 |
+ |
| 168 |
206 |
payload = { |
| 169 |
207 |
"experiment": "candidate_01_deferred_refinement", |
| 170 |
208 |
"author": "Simon-Pierre Boucher", |
| 179 |
217 |
"pure_q8": {"tokens_per_s": toks_per_s(q8_times), "fidelity_vs_q8_mean": 1.0}, |
| 180 |
218 |
}, |
| 181 |
219 |
"runs": results, |
|
220 |
+ "quality_bf16_judge": quality, |
| 182 |
221 |
} |
| 183 |
222 |
ts = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ") |
| 184 |
223 |
out_dir = REPO_ROOT / "results" / "candidate_01" / ts |