summaryrefslogtreecommitdiff
path: root/ep_run/casc_bp_train.py
diff options
context:
space:
mode:
Diffstat (limited to 'ep_run/casc_bp_train.py')
-rw-r--r--ep_run/casc_bp_train.py10
1 files changed, 8 insertions, 2 deletions
diff --git a/ep_run/casc_bp_train.py b/ep_run/casc_bp_train.py
index 5194c45..8047091 100644
--- a/ep_run/casc_bp_train.py
+++ b/ep_run/casc_bp_train.py
@@ -14,6 +14,8 @@ ap.add_argument('--lr', type=float, default=3e-4); ap.add_argument('--warmup', t
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('--wandb', default=''); ap.add_argument('--wandb_run', default='')
+ap.add_argument('--opt', choices=['adamw', 'muon'], 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 std (GPT-standard 0.02)
args = ap.parse_args()
torch.manual_seed(args.seed)
@@ -47,8 +49,12 @@ if args.tok_init > 0:
blocks = nn.ModuleList([Block(args.C, args.H) 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())
-opt = torch.optim.AdamW(params, lr=args.lr, weight_decay=1e-4)
-sched = torch.optim.lr_scheduler.LambdaLR(opt, lambda s: min(1.0, (s + 1) / max(args.warmup, 1)))
+if args.opt == 'muon':
+ from muon import build_hybrid
+ opt, sched = build_hybrid(blocks, params, args.lr, args.muon_lr, args.warmup)
+else:
+ opt = torch.optim.AdamW(params, lr=args.lr, weight_decay=1e-4)
+ sched = torch.optim.lr_scheduler.LambdaLR(opt, lambda s: min(1.0, (s + 1) / max(args.warmup, 1)))
def fwd(x):
z = tok(x) + pos(torch.arange(args.T, device=dev))[None]