summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorYuren Hao <yurenh2@illinois.edu>2026-07-10 02:56:29 -0500
committerYuren Hao <yurenh2@illinois.edu>2026-07-10 02:56:29 -0500
commit3d1bae3cd3e3333c2134669d2453fbb9d6a1adf4 (patch)
tree66597e28b70ca7e6338787a52608b7429d7b3cc5
parentf5861055f40a45805bc1d34f36fd1c005ad818b6 (diff)
Add --cosine (warmup->cosine decay to 0.1x lr) to both cascade trainers for full-epoch runs
-rw-r--r--ep_run/casc_bp_train.py11
-rw-r--r--ep_run/casc_eq_train.py11
2 files changed, 20 insertions, 2 deletions
diff --git a/ep_run/casc_bp_train.py b/ep_run/casc_bp_train.py
index 8047091..7ee84e8 100644
--- a/ep_run/casc_bp_train.py
+++ b/ep_run/casc_bp_train.py
@@ -17,6 +17,8 @@ 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)
+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)
args = ap.parse_args()
torch.manual_seed(args.seed)
dev = 'cuda' if torch.cuda.is_available() else 'cpu'
@@ -54,7 +56,14 @@ if args.opt == 'muon':
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)))
+ if args.cosine:
+ def _lrlam(s):
+ if s < args.warmup: return (s + 1) / max(args.warmup, 1)
+ p = min(1.0, (s - args.warmup) / max(1, args.steps - args.warmup))
+ return args.lr_min_ratio + 0.5 * (1 - args.lr_min_ratio) * (1 + math.cos(math.pi * p))
+ sched = torch.optim.lr_scheduler.LambdaLR(opt, _lrlam)
+ else:
+ 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]
diff --git a/ep_run/casc_eq_train.py b/ep_run/casc_eq_train.py
index 67ee98f..ccaa204 100644
--- a/ep_run/casc_eq_train.py
+++ b/ep_run/casc_eq_train.py
@@ -28,6 +28,8 @@ ap.add_argument('--compile', action='store_true') # torch.compile each blo
ap.add_argument('--sig_every', type=int, default=25) # tok-sigma refresh interval (amortized)
ap.add_argument('--beta_floor', type=float, default=0.0) # >0: floor beta_t (anti finite-beta SNR collapse at depth)
ap.add_argument('--beta_fixed', action='store_true') # disable sig^2 schedule, hold beta_t = args.beta constant
+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('--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('--gate_every', type=int, default=200) # in-training cos(EP,BP) telemetry; <=0 = fully BP-free (no bp_gate at all)
ap.add_argument('--gate_govern', action='store_true') # let gate cos adjust K/bscale (default: observe-only => training control is BP-free)
@@ -76,7 +78,14 @@ if args.opt == 'muon':
opt, sched = build_hybrid(blocks, all_params, args.lr, args.muon_lr, args.warmup)
else:
opt = torch.optim.AdamW(all_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.cosine:
+ def _lrlam(s):
+ if s < args.warmup: return (s + 1) / max(args.warmup, 1)
+ p = min(1.0, (s - args.warmup) / max(1, args.steps - args.warmup))
+ return args.lr_min_ratio + 0.5 * (1 - args.lr_min_ratio) * (1 + math.cos(math.pi * p))
+ sched = torch.optim.lr_scheduler.LambdaLR(opt, _lrlam)
+ else:
+ sched = torch.optim.lr_scheduler.LambdaLR(opt, lambda s: min(1.0, (s + 1) / max(args.warmup, 1)))
NBT = args.B * args.T
def free_states_graphed(x):