diff options
| author | CaptainJack2491 <jayrupnakawala@gmail.com> | 2026-08-29 23:35:10 +0100 |
|---|---|---|
| committer | CaptainJack2491 <jayrupnakawala@gmail.com> | 2026-08-29 23:35:10 +0100 |
| commit | fc56dd007142264ef029eb8c0f9db11a6cf3d298 (patch) | |
| tree | ce0a7e53c198b9d499363c32486327d3718e8bbe /src/models | |
| parent | 9a5912f7764dbc942b56f8114f2c878c1a6d44ff (diff) | |
fix(transformer): increase positional embedding and causal mask buffer to 1024 to support worst-case probe traces
Diffstat (limited to 'src/models')
| -rw-r--r-- | src/models/transformer.py | 4 |
1 files changed, 2 insertions, 2 deletions
diff --git a/src/models/transformer.py b/src/models/transformer.py index a128ad9..4eb690a 100644 --- a/src/models/transformer.py +++ b/src/models/transformer.py @@ -32,13 +32,13 @@ 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(256, d) # supports sequences up to 256 + self.pos = nn.Embedding(1024, d) # supports full probe traces up to 1024 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) self.register_buffer( "causal_triu", - torch.triu(torch.ones(256, 256, dtype=torch.bool), diagonal=1), + torch.triu(torch.ones(1024, 1024, dtype=torch.bool), diagonal=1), persistent=False, ) |
