""" 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()