| """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 |
| from eatmap.mae import MAEPretrainer, patchify |
| from eatmap.model import Encoder |
|
|
| 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) |
| flat = patchify(spec, 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())) |
| |
| 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 |
| from eatmap.probe import load_encoder |
| from eatmap.runner import export_weights |
|
|
| 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) |
|
|