From 7467b5c313393e9c01b5d877f45e680e0c35213f Mon Sep 17 00:00:00 2001 From: Yuren Hao Date: Wed, 29 Jul 2026 04:46:31 -0500 Subject: =?UTF-8?q?RESULT=2069:=20=E9=A5=B1=E5=92=8C=E7=9B=B2=E5=8C=BA?= =?UTF-8?q?=E8=A2=AB=E8=87=AA=E5=B7=B1=E7=9A=84=E6=B2=BB=E7=96=97=E6=B5=8B?= =?UTF-8?q?=E8=AF=95=E8=AF=81=E4=BC=AA(cap30/50=E5=85=B3=E9=97=AD0.0000,?= =?UTF-8?q?=E8=AF=81=E4=B9=A6cap=3D1=E7=88=86=E6=A8=A1=E5=9E=8B);=20cap?= =?UTF-8?q?=E4=BB=B7=E6=A0=BC=E2=89=880(=E7=94=A8=E6=88=B7=E6=8B=85?= =?UTF-8?q?=E5=BF=A7=E5=AE=9A=E4=BB=B7);=20=E5=AD=98=E6=B4=BB=E7=BA=A6?= =?UTF-8?q?=E6=9D=9F=3D=E9=A1=B6=E5=8D=8A+attn+1/=CE=B2+11=E9=A1=B9?= =?UTF-8?q?=E5=85=8D=E7=96=AB;=20dgain=E4=BD=8D=E7=A7=BB=E8=87=82=E5=9C=A8?= =?UTF-8?q?=E9=A3=9E(cos0.83=E7=9A=84=E4=BB=A3=E4=BB=B7=E7=94=B1CE?= =?UTF-8?q?=E8=A3=81=E5=86=B3)+mixffn=E8=A1=A5=E5=88=86=E8=A7=A3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_014FAPDWQ49M5Ye3NpTndTpn --- ep_run/casc_bp_train.py | 9 ++++++++- ep_run/casc_eq_train.py | 17 +++++++++++++++-- 2 files changed, 23 insertions(+), 3 deletions(-) (limited to 'ep_run') 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): diff --git a/ep_run/casc_eq_train.py b/ep_run/casc_eq_train.py index 07e2ad4..d7cfdf7 100644 --- a/ep_run/casc_eq_train.py +++ b/ep_run/casc_eq_train.py @@ -49,6 +49,12 @@ ap.add_argument('--bsign_rand', action='store_true') # random-sign beta per ste ap.add_argument('--bf16', action='store_true') # cast model to bf16 (E-accumulation + tok_sigma stay fp32) — the x0.5 cost lever, GATE before production ap.add_argument('--amp', action='store_true') # PROPER mixed precision: autocast(bf16) matmuls, fp32 params/states/d/E — amp_gate.py PASSED 2026-07-12 (cos 0.9682 vs fp32 0.9687); --bf16 naive-cast stays DEAD (state quantization, RESULT 11) ap.add_argument('--dtop_every', type=int, default=1) # 1 = exact (DEFAULT, BP-parity); 2 = fast mode (~20% cheaper, ~4% CE tax at high lr) +ap.add_argument('--dgain_top', type=float, default=1.0) # amplify d in STATE FORMATION for blocks + # >= L/2 (read cotangents stay true-d: 1st- + # order exact; unlocks 2nd-order response of + # threshold nonlinearities without global beta) +ap.add_argument('--dgain_all', type=float, default=1.0) # same, all blocks (displacement-vs-force probe) +ap.add_argument('--logit_cap', type=float, default=0.0) # >0: Gemma-2-style attn logit softcap ap.add_argument('--head_lr_mult', type=float, default=1.0) # W_out Adam-group LR multiplier (C768 # head-throttle compensation, RESULT 66) ap.add_argument('--res_gate', type=float, default=0.02) # legality residual threshold; 0.02 was calibrated @@ -258,7 +264,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): @@ -407,7 +419,8 @@ def relax(z0, zs, ins, outs, y, beta, K, x, bmask=None): else: i = prev.detach().requires_grad_(True) o = blocks[l](i, mask) - znew = o.detach().float() + d[l] + _dg = args.dgain_all * (args.dgain_top if l >= args.L // 2 else 1.0) + znew = o.detach().float() + (_dg * d[l] if _dg != 1.0 else d[l]) # damped (under-relaxed) mixing: geta<1 restores contraction on stiff operators # (wall-2 toolkit); fixed point unchanged (z = z + geta*(o+d-z) <=> z = o+d) mixed = znew if g_eff >= 1.0 else (zs[l] + g_eff * (znew - zs[l])) -- cgit v1.2.3