"""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)