summaryrefslogtreecommitdiff
path: root/src/models
diff options
context:
space:
mode:
Diffstat (limited to 'src/models')
-rw-r--r--src/models/transformer.py23
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: