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
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
|
"""Unit tests for Addendum 8: 6-arm token-space recurrence & decomposition."""
import pytest
import torch
import torch.nn as nn
from src.config import Config
from src.data import (
EOS_ID,
NOISE_SLOT_ID,
PAUSE_ID,
RAND_START_ID,
SEP_ID,
SYM_C,
SYM_COLON,
SYM_D,
SYM_EQ,
build_examples,
decode_tokens,
make_batch,
make_structured_trace,
next_prime,
sieve_primes,
)
from src.models.transformer import TransformerBaseline
def test_vocab_and_special_tokens():
cfg_none = Config(scratch_mode="none")
assert cfg_none.vocab == 11
assert cfg_none.eos_id == 10
assert cfg_none.pad_id == 11
cfg_scratch = Config(scratch_mode="filler")
assert cfg_scratch.vocab == 34
assert cfg_scratch.eos_id == 10
assert cfg_scratch.pad_id == 34
def test_make_structured_trace():
primes = sieve_primes(50)
# n=14 -> next prime 17. Candidates 15 (comp: d3:0), 16 (comp: d2:0), 17 (prime: d2:1, d3:1)
trace = make_structured_trace(14, 17, primes)
assert SYM_C in trace
assert SYM_EQ in trace
assert SYM_D in trace
assert SYM_COLON in trace
# n=2 -> next prime 3. Candidate 3 has no primes with p*p <= 3, trace should be [SYM_C, SYM_EQ, 3]
trace_2_3 = make_structured_trace(2, 3, primes)
assert trace_2_3 == [SYM_C, SYM_EQ, 3]
def test_decode_tokens_all_modes():
cfg_none = Config(scratch_mode="none")
assert decode_tokens([4, 3, EOS_ID], cfg_none) == 43
cfg_filler = Config(scratch_mode="filler")
# with separator #
assert decode_tokens([PAUSE_ID, PAUSE_ID, SEP_ID, 4, 3, EOS_ID], cfg_filler) == 43
# without separator (should fail / return -1)
assert decode_tokens([4, 3, EOS_ID], cfg_filler) == -1
def test_build_examples_and_loss_mask():
modes = ["none", "structured", "filler", "random_learned", "random_frozen", "random_noise"]
for mode in modes:
cfg = Config(scratch_mode=mode, range_start=2, range_end=30, scratch_len=16)
exs = build_examples([14], cfg)
assert len(exs) == 1
x, y = exs[0]
batch = make_batch(exs, cfg)
loss_mask = batch["loss_mask"][0]
if mode == "none":
assert decode_tokens(y, cfg) == 17
assert all(loss_mask[: len(y)])
elif mode == "structured":
assert SEP_ID in y
assert decode_tokens(y, cfg) == 17
assert all(loss_mask[: len(y)]) # structured trace is fully supervised
elif mode in ("filler", "random_learned", "random_frozen", "random_noise"):
assert SEP_ID in y
assert decode_tokens(y, cfg) == 17
# first 16 tokens must have loss_mask == 0
assert all(not m for m in loss_mask[:16])
# tokens after separator must have loss_mask == 1
assert all(m for m in loss_mask[16 : len(y)])
def test_random_tokens_no_digit_leakage():
for mode in ("random_learned", "random_frozen"):
cfg = Config(scratch_mode=mode, range_start=2, range_end=100, scratch_len=16)
exs = build_examples(list(range(2, 101)), cfg)
for _, y in exs:
random_segment = y[:16]
for tok in random_segment:
assert RAND_START_ID <= tok <= RAND_START_ID + 15
assert not (0 <= tok <= 9) # no digits leakage
def test_transformer_forward_all_arms():
modes = ["none", "structured", "filler", "random_learned", "random_frozen", "random_noise"]
for mode in modes:
cfg = Config(scratch_mode=mode, range_start=2, range_end=30, d_model=32, n_layers=1, n_heads=2)
exs = build_examples([14, 15, 16], cfg)
batch = make_batch(exs, cfg)
model = TransformerBaseline(cfg)
out = model(batch["x"], batch["y_in"])
logits = out["logits"]
assert logits.shape == (3, batch["y"].shape[1], cfg.vocab)
# compute masked loss
ce = nn.CrossEntropyLoss(reduction="none")
y_safe = batch["y"].clamp(max=cfg.vocab - 1)
loss_tokens = ce(logits.reshape(-1, cfg.vocab), y_safe.reshape(-1)).reshape(3, -1) * batch["loss_mask"].float()
loss = loss_tokens.sum() / batch["loss_mask"].sum().clamp(min=1)
assert not torch.isnan(loss)
loss.backward()
def test_random_frozen_zero_gradients():
cfg = Config(scratch_mode="random_frozen", range_start=2, range_end=30, d_model=32, n_layers=1, n_heads=2)
exs = build_examples([14, 15], cfg)
batch = make_batch(exs, cfg)
model = TransformerBaseline(cfg)
out = model(batch["x"], batch["y_in"])
logits = out["logits"]
loss = logits.sum()
loss.backward()
# random token embeddings (13..28) must have zero grad because they were detached
grad = model.embed.weight.grad
assert grad[RAND_START_ID : RAND_START_ID + 16].abs().sum() == 0
# non-random tokens must have non-zero grad
assert grad[:10].abs().sum() > 0
def test_random_noise_stochasticity():
cfg = Config(scratch_mode="random_noise", range_start=2, range_end=30, d_model=32, n_layers=1, n_heads=2)
exs = build_examples([14], cfg)
batch = make_batch(exs, cfg)
model = TransformerBaseline(cfg)
model.eval()
out1 = model(batch["x"], batch["y_in"])["logits"]
out2 = model(batch["x"], batch["y_in"])["logits"]
# Dynamic Gaussian noise should produce distinct logits between forward passes
assert not torch.allclose(out1, out2)
|