summaryrefslogtreecommitdiff
path: root/src/models/transformer.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/models/transformer.py')
-rw-r--r--src/models/transformer.py61
1 files changed, 61 insertions, 0 deletions
diff --git a/src/models/transformer.py b/src/models/transformer.py
new file mode 100644
index 0000000..32e158b
--- /dev/null
+++ b/src/models/transformer.py
@@ -0,0 +1,61 @@
+"""Transformer baseline: GPT-style causal decoder over [input digits | output digits].
+
+Fixed d_model (=128) matching the RNN arm. Parameter counts are logged per run but
+NOT gated to parity (design/preregistration.md: weight sharing is the studied variable).
+"""
+import torch
+import torch.nn as nn
+
+from src.config import Config
+from src.model_api import PrimeModel
+
+
+class CausalBlock(nn.Module):
+ def __init__(self, d: int, heads: int):
+ super().__init__()
+ self.ln1 = nn.LayerNorm(d)
+ self.ln2 = nn.LayerNorm(d)
+ self.attn = nn.MultiheadAttention(d, heads, batch_first=True)
+ self.mlp = nn.Sequential(nn.Linear(d, 4 * d), nn.GELU(), nn.Linear(4 * d, d))
+
+ def forward(self, x: torch.Tensor, attn_mask: torch.Tensor, key_pad: torch.Tensor) -> torch.Tensor:
+ a, _ = self.attn(self.ln1(x), self.ln1(x), self.ln1(x),
+ attn_mask=attn_mask, key_padding_mask=key_pad, need_weights=False)
+ x = x + a
+ x = x + self.mlp(self.ln2(x))
+ return x
+
+
+class TransformerBaseline(PrimeModel):
+ def __init__(self, cfg: Config):
+ super().__init__()
+ 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.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)
+
+ def forward(self, x: torch.Tensor, y_in: torch.Tensor) -> dict:
+ cfg = self.cfg
+ B, T_in = x.shape
+ T_out = y_in.shape[1]
+ 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)
+ # 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
+ key_pad = seq == cfg.pad_id # (B,T) True = ignore
+ h = e
+ for blk in self.blocks:
+ h = blk(h, blocked, key_pad)
+ # logits at position p predict token p+1 -> positions [T_in-1, T_in+T_out-2] predict y
+ logits = self.head(self.ln_f(h))[:, T_in - 1: T_in - 1 + T_out]
+ return {"logits": logits, "halt_steps": None}