From 044973e2a78e0faa1296ac370cefbfc2f6feadd6 Mon Sep 17 00:00:00 2001 From: Yuren Hao Date: Wed, 5 Aug 2026 06:11:06 -0500 Subject: =?UTF-8?q?=E4=BC=98=E5=8C=96=E5=99=A8=E4=BB=B7=E7=9B=AE=E8=A1=A8?= =?UTF-8?q?=E5=AE=9E=E8=A3=85:=20Lion/OLion/Adafactor2D/build=5Falt=20+=20?= =?UTF-8?q?--opt=20=E5=85=AD=E9=80=89=E9=A1=B9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit OLion按论文算法框(慢动量+Nesterov混合→NS→逐元素sign→γ=0.2定标+解耦wd; betas取Lion默认 0.9/0.99并注明为假设); Adafactor2D=行/列因子化二阶矩+RMS-1裁剪(07-11审计的"模拟原生Adam"); build_alt与build_hybrid同分割(矩阵走被测规则+其余AdamW), sgdm例外全参数SGD(全局部基线)。 四路CPU冒烟全过(损失下降, gate cos 1.0000)。 Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_014FAPDWQ49M5Ye3NpTndTpn --- ep_run/casc_eq_train.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) (limited to 'ep_run/casc_eq_train.py') diff --git a/ep_run/casc_eq_train.py b/ep_run/casc_eq_train.py index 05c1e4e..7ded9d9 100644 --- a/ep_run/casc_eq_train.py +++ b/ep_run/casc_eq_train.py @@ -22,7 +22,7 @@ ap.add_argument('--wandb_run', default='') ap.add_argument('--kmax', type=int, default=8) # adaptive fb rounds cap ap.add_argument('--noguard', action='store_true') # diagnosis: skip only non-finite grads ap.add_argument('--untie', action='store_true') # separate readout matrix (untied from tok) -ap.add_argument('--opt', choices=['adamw', 'muon'], default='adamw') +ap.add_argument('--opt', choices=['adamw', 'muon', 'sgdm', 'lion', 'olion', 'adafactor'], default='adamw') ap.add_argument('--muon_lr', type=float, default=0.02) ap.add_argument('--tok_init', type=float, default=0.0) # >0: init tok/pos with this std (GPT-standard 0.02) ap.add_argument('--compile', action='store_true') # torch.compile each block (free speed where supported) @@ -341,7 +341,12 @@ if args.bf16: if args.untie: with torch.no_grad(): W_out.data = W_out.data.to(torch.bfloat16) print('[bf16] model cast to bfloat16 (E-accum + sigma stay fp32)', flush=True) -if args.opt == 'muon': +if args.opt in ('sgdm', 'lion', 'olion', 'adafactor'): + from muon import build_alt + opt, sched = build_alt(args.opt, blocks, all_params, args.lr, args.warmup, + total_steps=(args.steps if args.cosine else 0), lr_min_ratio=args.lr_min_ratio, + lr_matrix=(args.muon_lr if args.muon_lr != 0.02 else None)) +elif args.opt == 'muon': from muon import build_hybrid opt, sched = build_hybrid(blocks, all_params, args.lr, args.muon_lr, args.warmup, muon_mom=args.muon_mom, adam_b1=args.adam_b1, -- cgit v1.2.3