diff options
| author | CaptainJack2491 <jayrupnakawala@gmail.com> | 2026-08-29 23:19:29 +0100 |
|---|---|---|
| committer | CaptainJack2491 <jayrupnakawala@gmail.com> | 2026-08-29 23:19:29 +0100 |
| commit | 7b0fddb02b089b82b8d12b2dcac17bf9527817da (patch) | |
| tree | 504a86a5f115f3b7bccd5cd716aab01415328dcc /src/models | |
| parent | 34734e1ab91253a33eddb1d6d77694a1367a7a9a (diff) | |
implement Addendum 8: 6-arm token-space recurrence suite (arms A, B, C, D1, D2, D3), tests, and jobs/e8.csv
Diffstat (limited to 'src/models')
| -rw-r--r-- | src/models/transformer.py | 23 |
1 files changed, 17 insertions, 6 deletions
diff --git a/src/models/transformer.py b/src/models/transformer.py index 32e158b..2b986ce 100644 --- a/src/models/transformer.py +++ b/src/models/transformer.py @@ -32,7 +32,7 @@ class TransformerBaseline(PrimeModel): self.cfg = cfg d = cfg.d_model self.embed = nn.Embedding(cfg.vocab + 1, d) # +1 row = pad - self.pos = nn.Embedding(64, d) # 3 in + 6 out max, generous + self.pos = nn.Embedding(256, d) # supports sequences up to 256 self.blocks = nn.ModuleList([CausalBlock(d, cfg.n_heads) for _ in range(cfg.n_layers)]) self.ln_f = nn.LayerNorm(d) self.head = nn.Linear(d, cfg.vocab) @@ -44,14 +44,25 @@ class TransformerBaseline(PrimeModel): seq = torch.cat([x, y_in], dim=1) # (B, T_in+T_out) T = seq.shape[1] device = seq.device - e = self.embed(seq) + self.pos(torch.arange(T, device=device)).unsqueeze(0) + + # Token embeddings + tok_embed = self.embed(seq) + if cfg.scratch_mode == "random_frozen": + # detach embeddings for tokens 'a'..'p' (IDs 13..28) + is_rand = (seq >= 13) & (seq <= 28) + tok_embed = torch.where(is_rand.unsqueeze(-1), tok_embed.detach(), tok_embed) + elif cfg.scratch_mode == "random_noise": + # inject fresh Gaussian noise at NOISE_SLOT_ID (ID 33) + is_noise = (seq == 33) + noise = torch.randn_like(tok_embed) + tok_embed = torch.where(is_noise.unsqueeze(-1), noise, tok_embed) + + e = tok_embed + self.pos(torch.arange(T, device=device)).unsqueeze(0) # causal mask: input positions (j < T_in) fully visible; output positions causal # (True = blocked, per torch.nn.MultiheadAttention bool convention) blocked = torch.zeros(T, T, dtype=torch.bool, device=device) - for i in range(T): - for j in range(T): - if j >= T_in and j > i: - blocked[i, j] = True + causal = torch.triu(torch.ones(T, T, dtype=torch.bool, device=device), diagonal=1) + blocked[:, T_in:] = causal[:, T_in:] key_pad = seq == cfg.pad_id # (B,T) True = ignore h = e for blk in self.blocks: |
