diff options
Diffstat (limited to 'ep_run/casc_eq_train.py')
| -rw-r--r-- | ep_run/casc_eq_train.py | 11 |
1 files changed, 10 insertions, 1 deletions
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): |
