diff options
Diffstat (limited to 'ep_run/casc_bp_train.py')
| -rw-r--r-- | ep_run/casc_bp_train.py | 35 |
1 files changed, 29 insertions, 6 deletions
diff --git a/ep_run/casc_bp_train.py b/ep_run/casc_bp_train.py index 7ee84e8..c056fed 100644 --- a/ep_run/casc_bp_train.py +++ b/ep_run/casc_bp_train.py @@ -19,6 +19,7 @@ ap.add_argument('--muon_lr', type=float, default=0.02) ap.add_argument('--tok_init', type=float, default=0.0) # >0: init tok/pos std (GPT-standard 0.02) ap.add_argument('--cosine', action='store_true') # warmup then cosine decay to lr_min_ratio*lr over --steps (long runs) ap.add_argument('--lr_min_ratio', type=float, default=0.1) +ap.add_argument('--qk_norm', action='store_true') # RMS-norm q,k per head before scores (OLMo2-style; bounds logits, analog-friendly) args = ap.parse_args() torch.manual_seed(args.seed) dev = 'cuda' if torch.cuda.is_available() else 'cpu' @@ -33,22 +34,44 @@ def get_batch(split): y = torch.stack([torch.from_numpy(data[i + 1:i + 1 + args.T].astype(np.int64)) for i in ix]) return x.to(dev), y.to(dev) +class CausalSelfAttn(nn.Module): + """explicit MHA (SDPA-backed) so we can QK-norm q,k per head before the scores.""" + def __init__(self, C, H, qk_norm=False): + super().__init__() + self.H, self.hd, self.qk_norm = H, C // H, qk_norm + self.qkv = nn.Linear(C, 3 * C) + self.proj = nn.Linear(C, C) + if qk_norm: + self.q_g = nn.Parameter(torch.ones(self.hd)) + self.k_g = nn.Parameter(torch.ones(self.hd)) + def forward(self, x): + B, T, C = x.shape + q, k, v = self.qkv(x).split(C, dim=2) + q = q.view(B, T, self.H, self.hd).transpose(1, 2) + k = k.view(B, T, self.H, self.hd).transpose(1, 2) + v = v.view(B, T, self.H, self.hd).transpose(1, 2) + if self.qk_norm: # RMS-norm over head_dim (OLMo2-style), learnable per-dim gain + q = q * torch.rsqrt(q.pow(2).mean(-1, keepdim=True) + 1e-6) * self.q_g + k = k * torch.rsqrt(k.pow(2).mean(-1, keepdim=True) + 1e-6) * self.k_g + 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 Block(nn.Module): - def __init__(self, C, H): + def __init__(self, C, H, qk_norm=False): super().__init__() self.ln1, self.ln2 = nn.LayerNorm(C), nn.LayerNorm(C) - self.attn = nn.MultiheadAttention(C, H, batch_first=True) + self.attn = CausalSelfAttn(C, H, qk_norm) self.ff = nn.Sequential(nn.Linear(C, 4 * C), nn.GELU(), nn.Linear(4 * C, C)) - def forward(self, z, mask): - h = self.ln1(z); a, _ = self.attn(h, h, h, attn_mask=mask, need_weights=False) - z = z + a; return z + self.ff(self.ln2(z)) + def forward(self, z, mask=None): + z = z + self.attn(self.ln1(z)) + return z + self.ff(self.ln2(z)) tok = nn.Embedding(vocab, args.C).to(dev) pos = nn.Embedding(args.T, args.C).to(dev) if args.tok_init > 0: with torch.no_grad(): tok.weight.normal_(0, args.tok_init); pos.weight.normal_(0, args.tok_init) -blocks = nn.ModuleList([Block(args.C, args.H) for _ in range(args.L)]).to(dev) +blocks = nn.ModuleList([Block(args.C, args.H, args.qk_norm) for _ in range(args.L)]).to(dev) mask = torch.triu(torch.full((args.T, args.T), float('-inf'), device=dev), 1) params = list(tok.parameters()) + list(pos.parameters()) + list(blocks.parameters()) if args.opt == 'muon': |
