maxact-fast / scripts /test_packing.py
ceselder's picture
pretrain: sequence packing (block-diag attn, per-marker injection) — no-leak + inject proofs pass; MFU bench pending
e6dc020
Raw
History Blame Contribute Delete
7.49 kB
"""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()