summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorCaptainJack2491 <jayrupnakawala@gmail.com>2026-08-29 23:32:30 +0100
committerCaptainJack2491 <jayrupnakawala@gmail.com>2026-08-29 23:32:30 +0100
commit9a5912f7764dbc942b56f8114f2c878c1a6d44ff (patch)
tree1c01d9335ffa8964e75a9206f5b439403aa3e830
parent7b0fddb02b089b82b8d12b2dcac17bf9527817da (diff)
optimize E8: precompute causal mask buffer, add eval early exit, and enable compile_model in e8.csv
-rw-r--r--jobs/e8.csv13
-rw-r--r--src/models/transformer.py8
-rw-r--r--src/train.py3
3 files changed, 16 insertions, 8 deletions
diff --git a/jobs/e8.csv b/jobs/e8.csv
index 0c393a0..153b373 100644
--- a/jobs/e8.csv
+++ b/jobs/e8.csv
@@ -1,7 +1,8 @@
job,model,seed,flags
-e8-armA-none,transformer,0,--range_end 1000 --max_steps 32 --device cuda --weight_decay 0.1 --max_train_steps 200000 --scratch_mode none
-e8-armB-structured,transformer,0,--range_end 1000 --max_steps 32 --device cuda --weight_decay 0.1 --max_train_steps 200000 --scratch_mode structured
-e8-armC-filler,transformer,0,--range_end 1000 --max_steps 32 --device cuda --weight_decay 0.1 --max_train_steps 200000 --scratch_mode filler --scratch_len 16
-e8-armD1-randlearn,transformer,0,--range_end 1000 --max_steps 32 --device cuda --weight_decay 0.1 --max_train_steps 200000 --scratch_mode random_learned --scratch_len 16
-e8-armD2-randfroz,transformer,0,--range_end 1000 --max_steps 32 --device cuda --weight_decay 0.1 --max_train_steps 200000 --scratch_mode random_frozen --scratch_len 16
-e8-armD3-randnoise,transformer,0,--range_end 1000 --max_steps 32 --device cuda --weight_decay 0.1 --max_train_steps 200000 --scratch_mode random_noise --scratch_len 16
+e8-armA-none,transformer,0,--range_end 1000 --max_steps 32 --device cuda --compile_model True --weight_decay 0.1 --max_train_steps 200000 --scratch_mode none
+e8-armB-structured,transformer,0,--range_end 1000 --max_steps 32 --device cuda --compile_model True --weight_decay 0.1 --max_train_steps 200000 --scratch_mode structured
+e8-armC-filler,transformer,0,--range_end 1000 --max_steps 32 --device cuda --compile_model True --weight_decay 0.1 --max_train_steps 200000 --scratch_mode filler --scratch_len 16
+e8-armD1-randlearn,transformer,0,--range_end 1000 --max_steps 32 --device cuda --compile_model True --weight_decay 0.1 --max_train_steps 200000 --scratch_mode random_learned --scratch_len 16
+e8-armD2-randfroz,transformer,0,--range_end 1000 --max_steps 32 --device cuda --compile_model True --weight_decay 0.1 --max_train_steps 200000 --scratch_mode random_frozen --scratch_len 16
+e8-armD3-randnoise,transformer,0,--range_end 1000 --max_steps 32 --device cuda --compile_model True --weight_decay 0.1 --max_train_steps 200000 --scratch_mode random_noise --scratch_len 16
+
diff --git a/src/models/transformer.py b/src/models/transformer.py
index 2b986ce..a128ad9 100644
--- a/src/models/transformer.py
+++ b/src/models/transformer.py
@@ -36,6 +36,11 @@ class TransformerBaseline(PrimeModel):
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),
+ persistent=False,
+ )
def forward(self, x: torch.Tensor, y_in: torch.Tensor) -> dict:
cfg = self.cfg
@@ -61,8 +66,7 @@ class TransformerBaseline(PrimeModel):
# 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)
- causal = torch.triu(torch.ones(T, T, dtype=torch.bool, device=device), diagonal=1)
- blocked[:, T_in:] = causal[:, T_in:]
+ blocked[:, T_in:] = self.causal_triu[:T, T_in:T]
key_pad = seq == cfg.pad_id # (B,T) True = ignore
h = e
for blk in self.blocks:
diff --git a/src/train.py b/src/train.py
index 2523325..8ce5394 100644
--- a/src/train.py
+++ b/src/train.py
@@ -105,6 +105,9 @@ def evaluate_gpu(model, gpu_ds: dict[str, torch.Tensor | int | bool], examples:
o = model(x, y_in)
nxt = o["logits"][:, -1].argmax(-1)
y_in = torch.cat([y_in, nxt[:, None]], dim=1)
+ # If all sequences in the batch have emitted at least one EOS, we can exit safely
+ if ((y_in[:, 1:] == cfg.eos_id).any(dim=1)).all():
+ break
gen = y_in[:, 1:].cpu()
for i, ex in enumerate(chunk_examples):