File size: 4,429 Bytes
fd3090c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
"""
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()