diff options
| author | Yuren Hao <yurenh2@illinois.edu> | 2026-07-29 04:46:31 -0500 |
|---|---|---|
| committer | Yuren Hao <yurenh2@illinois.edu> | 2026-07-29 04:46:31 -0500 |
| commit | 7467b5c313393e9c01b5d877f45e680e0c35213f (patch) | |
| tree | 8aa466c1b308170361c9ce353439dbc191dab2ff /ep_run/casc_bp_train.py | |
| parent | c67f91ba8f1eb70c64897641bb88cee4af249ed2 (diff) | |
RESULT 69: 饱和盲区被自己的治疗测试证伪(cap30/50关闭0.0000,证书cap=1爆模型); cap价格≈0(用户担忧定价); 存活约束=顶半+attn+1/β+11项免疫; dgain位移臂在飞(cos0.83的代价由CE裁决)+mixffn补分解
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014FAPDWQ49M5Ye3NpTndTpn
Diffstat (limited to 'ep_run/casc_bp_train.py')
| -rw-r--r-- | ep_run/casc_bp_train.py | 9 |
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): |
