From 9a5912f7764dbc942b56f8114f2c878c1a6d44ff Mon Sep 17 00:00:00 2001 From: CaptainJack2491 Date: Sat, 29 Aug 2026 23:32:30 +0100 Subject: optimize E8: precompute causal mask buffer, add eval early exit, and enable compile_model in e8.csv --- jobs/e8.csv | 13 +++++++------ src/models/transformer.py | 8 ++++++-- src/train.py | 3 +++ 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): -- cgit v1.2.3