""" Local shape/alignment test with a stub encoder. No GPU, no downloads. Verifies: - forward() output shape == (B, L_tgt, vocab) - generate() returns (B, 1 + n) - DecoderNoConditioning has the same trainable param count as the conditioned arm - the loss is finite and gradients reach every trainable tensor """ import torch import torch.nn as nn from model import DecoderNoConditioning, SemanticConditionedDecoder VOCAB = 131072 HIDDEN = 2048 D_MODEL = 1024 B, L_SRC, L_TGT = 2, 16, 12 class StubBackbone(nn.Module): def __init__(self): super().__init__() self.embedding = nn.Embedding(VOCAB, HIDDEN) def forward(self, input_ids, attention_mask=None): h = self.embedding(input_ids) return type("Out", (), {"last_hidden_state": h})() class StubEncoder(nn.Module): def __init__(self): super().__init__() self.config = type("C", (), {"hidden_size": HIDDEN})() self.backbone = StubBackbone() def forward(self, input_ids, attention_mask=None): return torch.randn(input_ids.shape[0], 256) class StubTokenizer: vocab_size = VOCAB bos_token_id = 0 eos_token_id = 1 pad_token_id = 2 def build(cls): """Build a model with a FROZEN stub encoder, mirroring the real setup.""" encoder = StubEncoder() for p in encoder.parameters(): p.requires_grad = False return cls(encoder=encoder, tokenizer=StubTokenizer(), d_model=D_MODEL, nhead=16, num_decoder_layers=1, dim_feedforward=1024, dropout=0.0) def main(): torch.manual_seed(0) tok = StubTokenizer() src = torch.randint(3, VOCAB, (B, L_SRC)) mask = torch.ones(B, L_SRC, dtype=torch.long) dec_in = torch.randint(3, VOCAB, (B, L_TGT)) labels = torch.randint(3, VOCAB, (B, L_TGT)) # --- conditioned arm --- model = build(SemanticConditionedDecoder) logits = model(input_ids=src, attention_mask=mask, decoder_input_ids=dec_in) assert logits.shape == (B, L_TGT, VOCAB), f"bad logits shape {logits.shape}" print("forward logits shape:", tuple(logits.shape), "OK") loss = nn.CrossEntropyLoss()(logits.reshape(-1, VOCAB), labels.reshape(-1)) assert torch.isfinite(loss), "loss is not finite" loss.backward() print("loss:", round(loss.item(), 4), "OK") no_grad = [n for n, p in model.named_parameters() if p.requires_grad and p.grad is None] assert not no_grad, f"trainable params got no gradient: {no_grad}" print("all trainable params received gradients OK") frozen_with_grad = [n for n, p in model.named_parameters() if not p.requires_grad and p.grad is not None] assert not frozen_with_grad, f"frozen params got gradients: {frozen_with_grad}" print("no gradients leaked into the frozen encoder OK") # --- generation --- model.eval() out = model.generate(src, mask, max_new_tokens=5, temperature=1.0) assert out.shape[0] == B and out.shape[1] >= 1, f"bad generate shape {out.shape}" assert out[0, 0].item() == tok.bos_token_id, "generation must start from BOS" print("generate shape:", tuple(out.shape), "starts with BOS OK") # --- ablation baseline: parameter parity --- a = build(SemanticConditionedDecoder) b = build(DecoderNoConditioning) pa = sum(p.numel() for p in a.parameters() if p.requires_grad) pb = sum(p.numel() for p in b.parameters() if p.requires_grad) print(f"trainable params: conditioned={pa:,} baseline={pb:,}") # The baseline carries exactly one extra parameter: `null_memory` # (d_model = 1024), the learned constant that replaces the input-derived # memory. That is 0.005% of the total, so the comparison stays fair. We # assert the difference is exactly that, rather than pretending it is zero. diff = pb - pa assert diff == D_MODEL, ( f"expected the baseline to differ by exactly d_model={D_MODEL} " f"(null_memory), got {diff}" ) print(f"parameter parity OK (baseline has +{diff} = null_memory, " f"{100*diff/pa:.4f}% of total)") # baseline forward logits_b = b(input_ids=src, attention_mask=mask, decoder_input_ids=dec_in) assert logits_b.shape == (B, L_TGT, VOCAB), f"bad baseline shape {logits_b.shape}" print("baseline forward shape:", tuple(logits_b.shape), "OK") print("\nALL CHECKS PASSED") if __name__ == "__main__": main()