diff options
| author | CaptainJack2491 <jayrupnakawala@gmail.com> | 2026-08-29 23:32:30 +0100 |
|---|---|---|
| committer | CaptainJack2491 <jayrupnakawala@gmail.com> | 2026-08-29 23:32:30 +0100 |
| commit | 9a5912f7764dbc942b56f8114f2c878c1a6d44ff (patch) | |
| tree | 1c01d9335ffa8964e75a9206f5b439403aa3e830 | |
| parent | 7b0fddb02b089b82b8d12b2dcac17bf9527817da (diff) | |
optimize E8: precompute causal mask buffer, add eval early exit, and enable compile_model in e8.csv
| -rw-r--r-- | jobs/e8.csv | 13 | ||||
| -rw-r--r-- | src/models/transformer.py | 8 | ||||
| -rw-r--r-- | 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): |
