summaryrefslogtreecommitdiff
path: root/ep_run/casc_eq_train.py
diff options
context:
space:
mode:
authorYuren Hao <yurenh2@illinois.edu>2026-08-05 06:11:06 -0500
committerYuren Hao <yurenh2@illinois.edu>2026-08-05 06:11:06 -0500
commit044973e2a78e0faa1296ac370cefbfc2f6feadd6 (patch)
treefb8fd3df6eaffb36e20757c5f28bce6046b8c607 /ep_run/casc_eq_train.py
parent98e145cda386167ce6d9d20120b4e265b3a6471a (diff)
优化器价目表实装: Lion/OLion/Adafactor2D/build_alt + --opt 六选项
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 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_014FAPDWQ49M5Ye3NpTndTpn
Diffstat (limited to 'ep_run/casc_eq_train.py')
-rw-r--r--ep_run/casc_eq_train.py9
1 files changed, 7 insertions, 2 deletions
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,