| """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 |
| from mxf.inject import get_layer, hooked, make_inject_hook, make_packed_inject_hook |
| from mxf.prompts import build_sft_ids |
|
|
| 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"]) |
| |
| |
| 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) |
|
|
| |
| 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" |
| |
| 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 |
|
|
| |
| 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" |
|
|
| |
| 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() |
|
|