mirror of
https://github.com/wassname/moral-maps.git
synced 2026-09-09 11:27:22 +08:00
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.
This commit is contained in:
@@ -0,0 +1,122 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user