ukung's picture
Add source code (encoder_loader, model, data, train, ablation, generate) + NOTES
fd3090c verified
Raw History Blame Contribute Delete
4.43 kB
"""
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()