traj-mc / code /dry_check.py
ttishere's picture
Publish code/dry_check.py
90ddd62 verified
Raw History Blame Contribute Delete
6.15 kB
"""
dry_check.py -- Step 3 dry-run verification of the BASE/OURS calibration.
Verifies, empirically, the four pre-registered invariants:
(a) OURS: per-sample measured mask ratio ~= its own t (inside a 4-sigma
binomial band), AND the mask rate is equal across the front/middle/back
thirds of the window (no systematic positional bias).
(b) BASE: total MASK count == 0 over all samples.
(c) BASE/OURS window hashes match item-by-item, and at every NON-mask position
the two arms are byte-identical.
(d) manifest carries dataset/split/counts; under pure C4 the CoT count == 0.
Also checks cross-run determinism: build_windows called twice with the same seed
yields identical hashes (so a separate BASE run and OURS run see the same windows).
"""
import os
import sys
import math
import argparse
import numpy as np
import torch
HERE = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, HERE)
import common as C
from calib.build_calib import build_windows, apply_noise, load_corpus
ok_all = True
report = {}
def main():
global ok_all
ap = argparse.ArgumentParser()
ap.add_argument("--n", type=int, default=32)
ap.add_argument("--seqlen", type=int, default=C.SEQLEN)
ap.add_argument("--seed", type=int, default=42)
ap.add_argument("--c4_split", type=str, default="train[:100000]")
ap.add_argument("--model_path", type=str, default=C.DEFAULT_MODEL_PATH)
args = ap.parse_args()
from transformers import AutoTokenizer
tok = AutoTokenizer.from_pretrained(args.model_path, trust_remote_code=True)
traindata, comp, _ = load_corpus("c4", args.c4_split, "train")
print(f"[dry] building {args.n} windows (seqlen={args.seqlen}, seed={args.seed})")
windows, hashes = build_windows(tok, traindata, args.n, args.seqlen, args.seed)
# cross-run determinism: same seed -> same windows
windows2, hashes2 = build_windows(tok, traindata, args.n, args.seqlen, args.seed)
det = bool(torch.equal(windows, windows2)) and hashes == hashes2
ok_all &= det
print(f"[det] rebuild with same seed identical: {det} -> {'PASS' if det else 'FAIL'}")
report["cross_run_determinism"] = det
# arms
base_ids = windows.clone() # BASE: t=0
ours_ids, t_list, measured, thirds = apply_noise(windows, args.seed) # OURS
# ── (b) BASE has zero MASK ────────────────────────────────────────────────
n_mask_base = int((base_ids == C.MASK_ID).sum())
b_ok = n_mask_base == 0
ok_all &= b_ok
print(f"[b] BASE total MASK count = {n_mask_base} (expect 0) "
f"-> {'PASS' if b_ok else 'FAIL'}")
report["base_mask_count"] = n_mask_base
# ── (a) OURS mask ratio within 4-sigma; no positional bias ────────────────
L = args.seqlen
inband = 0
worst_z = 0.0
third_devs = []
for i in range(args.n):
t = float(t_list[i])
m = float(measured[i])
sd = math.sqrt(max(t * (1 - t), 1e-12) / L)
z = abs(m - t) / max(sd, 1e-12)
worst_z = max(worst_z, z)
if abs(m - t) <= 4 * sd:
inband += 1
# positional bias: max |third_rate - overall| across the 3 thirds
third_devs.append(max(abs(x - m) for x in thirds[i]))
a1_ok = inband == args.n
# thirds: each third's rate must sit inside a 4-sigma band for a third-length window
sd_third = [math.sqrt(max(float(t_list[i]) * (1 - float(t_list[i])), 1e-12) / (L // 3))
for i in range(args.n)]
a2_ok = all(third_devs[i] <= 4 * sd_third[i] for i in range(args.n))
ok_all &= (a1_ok and a2_ok)
print(f"[a] OURS mask ratio in 4-sigma band: {inband}/{args.n} (worst z={worst_z:.2f}) "
f"-> {'PASS' if a1_ok else 'FAIL'}")
print(f"[a] OURS front/mid/back thirds within 4-sigma: "
f"{'PASS' if a2_ok else 'FAIL'} (max third deviation={max(third_devs):.4f})")
report["mask_ratio_inband"] = f"{inband}/{args.n}"
report["worst_z"] = worst_z
report["thirds_ok"] = a2_ok
report["t_mean"] = float(np.mean(t_list))
report["measured_mean"] = float(np.mean(measured))
# ── (c) hash identity + non-mask byte identity ────────────────────────────
h_base = [C.sha256_ids(base_ids[i]) for i in range(args.n)]
h_pre = hashes
c1_ok = h_base == h_pre
nonmask = ours_ids != C.MASK_ID
c2_ok = bool(torch.equal(ours_ids[nonmask], base_ids[nonmask]))
# and OURS differs from BASE exactly at the masked positions
changed = (ours_ids != base_ids)
c3_ok = bool(torch.equal(changed, (ours_ids == C.MASK_ID) & (base_ids != C.MASK_ID)))
ok_all &= (c1_ok and c2_ok and c3_ok)
print(f"[c] window hashes BASE==pre-noise item-by-item: "
f"{'PASS' if c1_ok else 'FAIL'}")
print(f"[c] non-mask positions byte-identical BASE vs OURS: "
f"{'PASS' if c2_ok else 'FAIL'} "
f"({int(nonmask.sum())} positions compared)")
print(f"[c] OURS differs from BASE ONLY at masked positions: "
f"{'PASS' if c3_ok else 'FAIL'}")
report["hash_identity"] = c1_ok
report["nonmask_byte_identity"] = c2_ok
report["diff_only_at_mask"] = c3_ok
# ── (d) corpus composition ────────────────────────────────────────────────
cot_count = comp["cot"]["count"]
d_ok = (cot_count == 0)
ok_all &= d_ok
print(f"[d] corpus: c4={comp['c4']['dataset']} split={comp['c4']['split']} "
f"n_docs={comp['c4']['count']}; CoT count={cot_count} (expect 0) "
f"-> {'PASS' if d_ok else 'FAIL'}")
report["corpus_composition"] = comp
report["pass"] = ok_all
C.dump_json(report, os.path.join(HERE, "results", "calib",
f"{C.git_hash()}_dry_check.json"))
print(f"\n=== DRY CHECK {'ALL PASS' if ok_all else 'FAIL'} ===")
sys.exit(0 if ok_all else 1)
if __name__ == "__main__":
main()