From 35a9228dde348705040e4149f4da2f59fd37b9a8 Mon Sep 17 00:00:00 2001 From: Yuren Hao Date: Fri, 10 Jul 2026 08:00:21 -0500 Subject: Stage-1 epoch blowup @12100: sig story REFUTED (only +8%, cos fine till after); leading indicator=drift-guard skips -> contractivity bifurcation in nudged relaxation (cascade Hopf wall); add resume/sig0/final_ln; A/B/C diagnostic launched --- ep_run/casc_bp_train.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) (limited to 'ep_run/casc_bp_train.py') diff --git a/ep_run/casc_bp_train.py b/ep_run/casc_bp_train.py index c056fed..7647fae 100644 --- a/ep_run/casc_bp_train.py +++ b/ep_run/casc_bp_train.py @@ -20,6 +20,7 @@ ap.add_argument('--tok_init', type=float, default=0.0) # >0: init tok/pos std ( 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) +ap.add_argument('--final_ln', action='store_true') # final LayerNorm before readout (standard GPT; bounds sig_tok growth -> keeps beta/estimator healthy on long runs) args = ap.parse_args() torch.manual_seed(args.seed) dev = 'cuda' if torch.cuda.is_available() else 'cpu' @@ -73,7 +74,8 @@ if args.tok_init > 0: tok.weight.normal_(0, args.tok_init); pos.weight.normal_(0, args.tok_init) 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()) +ln_f = nn.LayerNorm(args.C).to(dev) if args.final_ln else nn.Identity() +params = list(tok.parameters()) + list(pos.parameters()) + list(blocks.parameters()) + list(ln_f.parameters()) if args.opt == 'muon': from muon import build_hybrid opt, sched = build_hybrid(blocks, params, args.lr, args.muon_lr, args.warmup) @@ -91,7 +93,7 @@ else: def fwd(x): z = tok(x) + pos(torch.arange(args.T, device=dev))[None] for b in blocks: z = b(z, mask) - return z @ tok.weight.t() + return ln_f(z) @ tok.weight.t() @torch.no_grad() def evaluate(nb=6): -- cgit v1.2.3