Files
moral-maps/scripts/smoke_batch_parity.py
T
wassname e996d57051 batch the eval: guided_rollout_batch + 4× speedup
Sequential eval was the bottleneck (~12s/vignette × 131 = 27 min/pass; with
bidirectional ±C × 14 methods, that projected to ~17 h). Three model calls
per row (phase1 generate, scoring forward, cosmetic continuation) became one
phase1 + one scoring per *batch*; continuation generate dropped (callers
only use p_true + pmass_format).

Parity smoke (scripts/smoke_batch_parity.py): float32 is bit-exact (max
Δp_true=0.0000 over 16 prompts, 4.3× speedup at limit=4). bf16 drifts on
individual rows — greedy argmax flips at near-tie tokens then phase1
diverges — but pmass agrees within 0.03 (scoring forward correct) and
aggregates over 131 vignettes will average out the per-row noise.

Other changes shipped in this commit:
- core.py: analyse() now returns raw_pmass dict alongside raw p_true (callers
  needed per-(vid,cond,frame) pmass for diagnostic warnings).
- guided.py guided_rollout: warn + log top-5 when pmass<0.9 (catches OOD
  steering / format-broken vignettes without a separate audit pass).
- eval.py: pre-tokenize a sample to log expected prompt+cache budget so OOM
  is predictable from the SHOULD line; group items by frame so each batch
  shares schema_hint + prefill.
2026-05-03 06:42:17 +08:00

123 lines
4.7 KiB
Python

"""Parity smoke: guided_rollout vs guided_rollout_batch on a small vignette subset.
Asserts p_true and pmass_format match within fp tolerance. Same chat template,
same prompts, same model, same generation kwargs -- only batching differs.
usage:
uv run python scripts/smoke_batch_parity.py --model Qwen/Qwen3-0.6B --limit 4
"""
from __future__ import annotations
import argparse
import time
import torch
from loguru import logger
from transformers import AutoModelForCausalLM, AutoTokenizer
from tinymfv.core import CONDITIONS, FRAMES
from tinymfv.data import load_vignettes
from tinymfv.guided import guided_rollout, guided_rollout_batch, choice_token_ids_tf
def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--model", default="Qwen/Qwen3-0.6B")
ap.add_argument("--limit", type=int, default=4)
ap.add_argument("--max-think-tokens", type=int, default=32)
ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
ap.add_argument("--dtype", default="bfloat16")
args = ap.parse_args()
rows = load_vignettes("")[: args.limit]
logger.info(f"{len(rows)} vignettes; testing parity")
dtype = getattr(torch, args.dtype)
tok = AutoTokenizer.from_pretrained(args.model)
if tok.pad_token is None:
tok.pad_token = tok.eos_token
tok.padding_side = "left"
model = AutoModelForCausalLM.from_pretrained(args.model, dtype=dtype).to(args.device).eval()
choice_ids = choice_token_ids_tf(tok)
# --- Sequential ---
t0 = time.time()
seq_results = [] # list of (vid, cond, frame, p_true, pmass)
for r in rows:
for cond in CONDITIONS:
for frame, fr in FRAMES.items():
res = guided_rollout(
model, tok,
user_prompt=r[cond],
choice_token_ids=choice_ids,
max_think_tokens=args.max_think_tokens,
schema_hint=fr["q"],
prefill=fr["prefill"],
)
seq_results.append((r["id"], cond, frame, res.p_true, res.pmass_format))
seq_elapsed = time.time() - t0
logger.info(f"sequential: {seq_elapsed:.1f}s ({len(seq_results)} prompts)")
# --- Batched ---
t0 = time.time()
batch_results = []
for frame, fr in FRAMES.items():
for cond in CONDITIONS:
user_prompts = [r[cond] for r in rows]
outs = guided_rollout_batch(
model, tok,
user_prompts=user_prompts,
choice_token_ids=choice_ids,
max_think_tokens=args.max_think_tokens,
schema_hint=fr["q"],
prefill=fr["prefill"],
)
for r, o in zip(rows, outs):
batch_results.append((r["id"], cond, frame, o.p_true, o.pmass_format))
batch_elapsed = time.time() - t0
logger.info(f"batched: {batch_elapsed:.1f}s (speedup={seq_elapsed/batch_elapsed:.1f}x)")
# --- Compare ---
seq_d = {(vid, c, f): (pt, pm) for vid, c, f, pt, pm in seq_results}
batch_d = {(vid, c, f): (pt, pm) for vid, c, f, pt, pm in batch_results}
assert set(seq_d) == set(batch_d), "key mismatch"
n = 0
max_pt_diff, max_pm_diff = 0.0, 0.0
rows_out = []
for k in seq_d:
spt, spm = seq_d[k]
bpt, bpm = batch_d[k]
d_pt = abs(spt - bpt)
d_pm = abs(spm - bpm)
max_pt_diff = max(max_pt_diff, d_pt)
max_pm_diff = max(max_pm_diff, d_pm)
rows_out.append((k, spt, bpt, d_pt, spm, bpm, d_pm))
n += 1
from tabulate import tabulate
print()
print(tabulate(
[(f"{k[0][:8]}|{k[1]}|{k[2]}", spt, bpt, d_pt, spm, bpm, d_pm)
for (k, spt, bpt, d_pt, spm, bpm, d_pm) in rows_out],
headers=["key", "p_true_seq", "p_true_bat", "Δp_true", "pm_seq", "pm_bat", "Δpm"],
floatfmt="+.4f", tablefmt="tsv",
))
# bf16 batched greedy decoding can pick different argmax than per-row greedy
# when two tokens tie within bf16 precision. The phase1 think rollout then
# diverges and per-row p_true drifts. float32 is bit-exact (use --dtype float32
# to verify the batching logic itself). At aggregate eval (131 vignettes
# averaged) the bf16 drift averages out; we accept it.
TOL = 0.20 if args.dtype != "float32" else 0.001
cue = "🟢" if (max_pt_diff < TOL and max_pm_diff < TOL) else "🔴"
print(f"\n{cue} max Δp_true={max_pt_diff:.4f} max Δpmass={max_pm_diff:.4f} (tol={TOL})")
print(f"speedup: {seq_elapsed/batch_elapsed:.1f}x ({len(seq_results)} prompts)")
if max_pt_diff >= TOL or max_pm_diff >= TOL:
raise SystemExit(f"PARITY FAILED: Δp_true={max_pt_diff:.4f} Δpmass={max_pm_diff:.4f}")
if __name__ == "__main__":
main()