summaryrefslogtreecommitdiff
path: root/ep_run/casc_bp_train.py
diff options
context:
space:
mode:
Diffstat (limited to 'ep_run/casc_bp_train.py')
-rw-r--r--ep_run/casc_bp_train.py9
1 files changed, 8 insertions, 1 deletions
diff --git a/ep_run/casc_bp_train.py b/ep_run/casc_bp_train.py
index 42acf16..4ecfee9 100644
--- a/ep_run/casc_bp_train.py
+++ b/ep_run/casc_bp_train.py
@@ -13,6 +13,7 @@ ap.add_argument('--B', type=int, default=24); ap.add_argument('--steps', type=in
ap.add_argument('--lr', type=float, default=3e-4); ap.add_argument('--warmup', type=int, default=200)
ap.add_argument('--seed', type=int, default=0)
ap.add_argument('--save_every', type=int, default=500); ap.add_argument('--log', type=int, default=200)
+ap.add_argument('--logit_cap', type=float, default=0.0) # >0: Gemma-2-style attn logit softcap
ap.add_argument('--watch_every', type=int, default=2000) # wandb-only telemetry: weight/act RMS
ap.add_argument('--wandb', default='auto') # ON BY DEFAULT; 'auto' = per-regime project (ept-fineweb-72m / ept-tinystories-42m); --wandb '' to disable
ap.add_argument('--wandb_run', default='')
@@ -114,7 +115,13 @@ class Olmo2Attn(nn.Module):
q = self.rope(q.view(B, T, self.H, self.hd).transpose(1, 2))
k = self.rope(k.view(B, T, self.H, self.hd).transpose(1, 2))
v = v.view(B, T, self.H, self.hd).transpose(1, 2)
- y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
+ if args.logit_cap > 0:
+ lg = (q @ k.transpose(-2, -1)) * (self.hd ** -0.5)
+ lg = args.logit_cap * torch.tanh(lg / args.logit_cap)
+ cm = torch.ones(T, T, dtype=torch.bool, device=x.device).tril()
+ y = lg.masked_fill(~cm, float('-inf')).softmax(-1) @ v
+ else:
+ y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
return self.proj(y.transpose(1, 2).contiguous().view(B, T, C))
class Olmo2Block(nn.Module):