"""Correctness proofs for packed pretraining (run on the training box, 1 GPU): (a) no cross-example leak — perturb an EARLIER example's tokens+vector; every LATER example's logits must be bit-identical (block-diag mask). (Also the task-literal direction: perturb a later example, check an earlier one — trivially guaranteed by causality.) (b) per-marker injection — swap ONE example's probe vector; its target-region logits must change materially, all other examples' logits identical. (+) packed vs unpacked parity — each example's packed logits vs the same example run alone through the legacy path (2D mask, default positions). python scripts/test_packing.py --data-dir /workspace/mxf/data/smoke_pre """ import argparse import copy import importlib.util import json import os import sys import numpy as np import torch from transformers import AutoModelForCausalLM, AutoTokenizer sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src")) from mxf.config import D_MODEL, INJECT_LAYER, MODEL, STEER_COEFF # noqa: E402 from mxf.inject import get_layer, hooked, make_inject_hook, make_packed_inject_hook # noqa: E402 from mxf.prompts import build_sft_ids # noqa: E402 spec = importlib.util.spec_from_file_location( "pretrain", os.path.join(os.path.dirname(__file__), "pretrain.py")) pretrain = importlib.util.module_from_spec(spec) spec.loader.exec_module(pretrain) PACK_LEN = 768 MAX_SEQ = 192 def main(): ap = argparse.ArgumentParser() ap.add_argument("--data-dir", default="/workspace/mxf/data/smoke_pre") a = ap.parse_args() dev = "cuda:0" tok = AutoTokenizer.from_pretrained(MODEL) if tok.pad_token is None: tok.pad_token = tok.eos_token records = [json.loads(l) for l in open(f"{a.data_dir}/records.jsonl")][:8] n_vecs = os.path.getsize(f"{a.data_dir}/vecs.f32") // (4 * D_MODEL) vecs = np.memmap(f"{a.data_dir}/vecs.f32", dtype=np.float32, mode="r", shape=(n_vecs, D_MODEL)) model = AutoModelForCausalLM.from_pretrained(MODEL, torch_dtype=torch.bfloat16, attn_implementation="sdpa", device_map={"": dev}) model.eval() layer = get_layer(model, INJECT_LAYER) toks = [] for r in records[:3]: ids, labs, pos, vidx = *build_sft_ids(tok, r["target_text"]), r["vec_idx"] toks.append((ids[:MAX_SEQ], labs[:MAX_SEQ], pos, vidx)) blk = pretrain.pack_examples(toks, PACK_LEN, seed=0) assert len(blk) == 1, f"expected 1 block, got {len(blk)}" blk = blk[0] K = len(blk["seg_lens"]) starts = np.concatenate([[0], np.cumsum(blk["seg_lens"])]).astype(int) spans = [slice(starts[j], starts[j + 1]) for j in range(K)] print(f"block: {K} examples, seg_lens={blk['seg_lens']}, markers={blk['markers']}, " f"vec_idxs={blk['vec_idxs']}, total={starts[-1]}/{PACK_LEN}") causal = torch.tril(torch.ones(PACK_LEN, PACK_LEN, dtype=torch.bool, device=dev)) def fwd(block, vec_idxs): ii, ll, pi, seg, rows, cols, _ = pretrain.pack_batch([block], PACK_LEN, tok.pad_token_id) assert (ii[0, cols] == ii[0, cols[0]]).all(), "marker positions misaligned" mask4 = pretrain.packed_attn_mask(seg.to(dev), causal, torch.bfloat16) vmat = torch.from_numpy(np.asarray(vecs[list(vec_idxs)])) hook = make_packed_inject_hook(vmat, rows, cols, STEER_COEFF, dev, torch.bfloat16) with hooked(layer, hook), torch.no_grad(): out = model(input_ids=ii.to(dev), attention_mask=mask4, position_ids=pi.to(dev), use_cache=False) return out.logits.float().cpu()[0] def perturb_tokens(block, j, seed): """Replace example j's TARGET tokens (same length, prompt+marker untouched).""" b = copy.deepcopy(block) rng = np.random.default_rng(seed) for t in range(spans[j].start, spans[j].stop): if b["labels"][t] != -100: b["ids"][t] = int(rng.integers(1000, 30000)) b["labels"][t] = b["ids"][t] return b base_v = list(blk["vec_idxs"]) # the smoke vec bank contains duplicate rows (same probe reused across records) — pick a swap # direction genuinely different from every direction used in the block bank = np.asarray(vecs[: min(200, len(vecs))], dtype=np.float64) bank /= np.linalg.norm(bank, axis=1, keepdims=True) used = bank[base_v] swap = int(np.argmin(np.abs(bank @ used.T).max(axis=1))) print(f"swap direction: row {swap}, max |cos| to block's directions = " f"{np.abs(bank[swap] @ used.T).max():.3f}") L0 = fwd(blk, base_v) # ---- (a) no cross-example leak: perturb FIRST example (tokens + vector) ---- v2 = list(base_v); v2[0] = swap L1 = fwd(perturb_tokens(blk, 0, seed=1), v2) print("\n(a) perturb example 0 (tokens+vec):") print(f" example 0 span logits max|Δ| = {(L1[spans[0]] - L0[spans[0]]).abs().max():.4f} (sanity: should be LARGE)") for j in range(1, K): d = (L1[spans[j]] - L0[spans[j]]).abs().max().item() print(f" example {j} span logits max|Δ| = {d:.2e} (must be < 1e-3)") assert d < 1e-3, "CROSS-EXAMPLE LEAK — block-diagonal mask broken" # task-literal direction: perturb example 1, check example 0 v3 = list(base_v); v3[1] = swap L2 = fwd(perturb_tokens(blk, 1, seed=2), v3) d = (L2[spans[0]] - L0[spans[0]]).abs().max().item() print(f" [literal direction] perturb example 1 → example 0 max|Δ| = {d:.2e} (must be < 1e-3)") assert d < 1e-3 # ---- (b) per-marker injection: swap ONLY example 1's vector ---- v4 = list(base_v); v4[1] = swap L3 = fwd(blk, v4) tgt1 = [t for t in range(spans[1].start, spans[1].stop) if blk["labels"][t] != -100] d_tgt = (L3[tgt1] - L0[tgt1]).abs() print("\n(b) swap example 1's probe vector only:") print(f" example 1 target-region logits: max|Δ| = {d_tgt.max():.3f}, mean|Δ| = {d_tgt.mean():.4f} (must be material)") for j in [0, 2]: d = (L3[spans[j]] - L0[spans[j]]).abs().max().item() print(f" example {j} span logits max|Δ| = {d:.2e} (must be < 1e-3)") assert d < 1e-3 assert d_tgt.max() > 0.5, "vector swap had no material effect — injection not firing per marker" # ---- (+) packed vs unpacked parity, per example ---- print("\n(+) packed vs unpacked (legacy path) parity:") by_vidx = {t[3]: t for t in toks} for j in range(K): ids, labs, pos, vidx = by_vidx[blk["vec_idxs"][j]] ii = torch.tensor([ids]); attn = torch.ones_like(ii, dtype=torch.bool) v = torch.from_numpy(np.asarray(vecs[vidx])).unsqueeze(0) hook = make_inject_hook([v], [pos], STEER_COEFF, dev, torch.bfloat16) with hooked(layer, hook), torch.no_grad(): lu = model(input_ids=ii.to(dev), attention_mask=attn.to(dev)).logits.float().cpu()[0] lp = L0[spans[j]] d = (lu - lp).abs() top_match = (lu.argmax(-1) == lp.argmax(-1)).float().mean().item() print(f" example {j}: max|Δ| = {d.max():.4f}, mean|Δ| = {d.mean():.5f}, " f"top-1 agreement = {top_match:.1%} (bf16 kernel noise expected)") print("\nALL PACKING TESTS PASSED") if __name__ == "__main__": main()