mae-cross-objective / scripts /smoke_mae.py
quinnlue's picture
MAE cross-objective mixture-ranking campaign
ffdcfe7 verified
Raw
History Blame Contribute Delete
5.06 kB
"""Correctness gates for the MAE graft. Run before any GPU time is spent.
Four things have to hold, and each has a failure mode that would otherwise show
up as a plausible loss curve and a worthless encoder:
1. patchify's token order equals the encoder's Conv2d token order. Asserted as
an exact algebraic identity, not a shape check.
2. The masker's mask/keep/restore indices are mutually consistent.
3. A forward pass produces a finite loss and gradients reach the student.
4. An MAE export loads through probe.load_encoder and accepts LoRA injection --
i.e. the readout really is objective-agnostic, which is the premise the whole
cross-objective comparison rests on.
"""
import sys
import tempfile
from pathlib import Path
import torch
sys.path.insert(0, "/workspace/code/eat-map-regmix")
from eatmap.config import RunConfig, load_config # noqa: E402
from eatmap.mae import MAEPretrainer, patchify # noqa: E402
from eatmap.model import Encoder # noqa: E402
FAIL = []
def check(label, cond, detail=""):
print(f" {'PASS' if cond else 'FAIL'} {label}{' ' + detail if detail else ''}")
if not cond:
FAIL.append(label)
print("\n[1] patchify order == Encoder.patches order")
torch.manual_seed(0)
for (pt, pf) in [(16, 16), (4, 64)]:
frames, mels = 1024, 64
grid = (frames // pt, mels // pf)
enc = Encoder(32, 1, 4, 4, pt, pf, grid)
spec = torch.randn(2, 1, frames, mels)
conv_tokens = enc.patches(spec) # (B, N, D)
flat = patchify(spec, pt, pf) # (B, N, pt*pf)
linear = flat @ enc.patch_embed.weight.flatten(1).T + enc.patch_embed.bias
err = (conv_tokens - linear).abs().max().item()
check(f"patch {pt}x{pf} grid {grid}", err < 1e-4, f"max|conv-patchify@W| = {err:.2e}")
print("\n[2] masker index consistency")
cfg_path = Path("/workspace/code/eat-map-regmix/configs")
config = load_config([cfg_path / "base.yaml", cfg_path / "scale_15m.yaml",
cfg_path / "budget_proxy.yaml", Path("/workspace/configs/mae.yaml")])
model = MAEPretrainer(config)
keep, restore, mask = model.masker(8, torch.device("cpu"))
n = config.num_patches
check("visible count", keep.shape[1] == config.visible_patches,
f"{keep.shape[1]} == {config.visible_patches} of {n}")
check("mask marks exactly the non-kept", int(mask.sum(1)[0]) == n - config.visible_patches,
f"{int(mask.sum(1)[0])} masked")
check("kept positions are unmasked", bool((~mask.gather(1, keep)).all()))
# restoring [kept ; masked] must recover canonical order
probe_ids = torch.arange(n).unsqueeze(0).expand(8, -1)
shuffled = torch.cat([keep, torch.nonzero(mask[0]).squeeze(-1).unsqueeze(0).expand(8, -1)], 1)
check("ids_restore inverts the shuffle", bool((shuffled.gather(1, restore)[0].sort().values
== probe_ids[0]).all()))
print("\n[3] forward + backward")
spec = torch.randn(4, 1, config.model.target_frames, config.model.mel_bins)
keep, restore, mask = model.masker(4 * config.objective.clone_batch, torch.device("cpu"))
loss, frame, utt = model(spec, keep, restore, mask)
check("loss is finite", torch.isfinite(loss).item(), f"loss = {loss.item():.4f}")
check("utterance term is zero", float(utt) == 0.0)
loss.backward()
g = model.student.blocks[0].attn.qkv.weight.grad
check("gradient reaches student block 0", g is not None and torch.isfinite(g).all().item())
counts = {"encoder": sum(p.numel() for p in model.student.parameters()),
"decoder": sum(p.numel() for p in model.decoder.parameters())}
print(f" encoder {counts['encoder']/1e6:.2f}M decoder {counts['decoder']/1e6:.2f}M")
print("\n[4] export -> probe.load_encoder -> LoRA injection")
from eatmap.lora import LoRAConfig, inject_lora # noqa: E402
from eatmap.probe import load_encoder # noqa: E402
from eatmap.runner import export_weights # noqa: E402
with tempfile.TemporaryDirectory() as tmp:
out = Path(tmp) / "export"
export_weights(model, config, out, step=1)
from safetensors.torch import load_file
keys = load_file(str(out / "model.safetensors")).keys()
check("export has no teacher tensors", not any(k.startswith("teacher.") for k in keys))
check("export has student tensors", any(k.startswith("student.") for k in keys))
enc, loaded_cfg = load_encoder(out, torch.device("cpu"))
check("probe.load_encoder accepts MAE export", enc is not None)
ref = model.student.blocks[0].attn.qkv.weight
check("loaded weights match the trained student",
torch.equal(enc.blocks[0].attn.qkv.weight, ref))
n_wrapped = inject_lora(enc, LoRAConfig())
check("LoRA wraps the expected module count", n_wrapped == 4 * config.model.depth,
f"{n_wrapped} == 4 x {config.model.depth}")
cls, patches = enc(spec)
check("encoder forward returns (cls, patches)",
cls.shape == (4, config.model.embed_dim) and patches.shape[1] == config.num_patches)
print("\n" + ("ALL GATES PASSED" if not FAIL else f"FAILED: {FAIL}"))
sys.exit(1 if FAIL else 0)