diff options
Diffstat (limited to 'ep_run/casc_bp_train.py')
| -rw-r--r-- | ep_run/casc_bp_train.py | 9 |
1 files changed, 7 insertions, 2 deletions
diff --git a/ep_run/casc_bp_train.py b/ep_run/casc_bp_train.py index 0515ec1..42acf16 100644 --- a/ep_run/casc_bp_train.py +++ b/ep_run/casc_bp_train.py @@ -13,6 +13,7 @@ ap.add_argument('--B', type=int, default=24); ap.add_argument('--steps', type=in ap.add_argument('--lr', type=float, default=3e-4); ap.add_argument('--warmup', type=int, default=200) 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('--watch_every', type=int, default=2000) # wandb-only telemetry: weight/act RMS ap.add_argument('--wandb', default='auto') # ON BY DEFAULT; 'auto' = per-regime project (ept-fineweb-72m / ept-tinystories-42m); --wandb '' to disable ap.add_argument('--wandb_run', default='') ap.add_argument('--amp', action='store_true') # bf16 autocast fwd/loss (no scaler needed for bf16) @@ -219,7 +220,7 @@ for step in range(start_step, args.steps + 1): if args.qcomp_bits > 0: with torch.no_grad(): for p, q in zip(params, QSAVE): p.copy_(q) - torch.nn.utils.clip_grad_norm_(params, 1.0) + CLIP_NORM = float(torch.nn.utils.clip_grad_norm_(params, 1.0)) opt.step(); sched.step() if args.qup_bits > 0: with torch.no_grad(): @@ -236,7 +237,11 @@ for step in range(start_step, args.steps + 1): print(f'step {step:5d}/{args.steps} | train {loss.item():.4f} val {val:.4f} (best {best:.4f}) ' f'| {step/max(time.time()-t0,1e-9):.2f} it/s', flush=True) if wb is not None: - try: wb.log({'train_ce': loss.item(), 'val_ce': val, 'best': best}, step=step) + _aux = {'clip_norm': CLIP_NORM, 'clip_fired': float(CLIP_NORM > 1.0), 'gn': CLIP_NORM} + if args.watch_every > 0 and step % args.watch_every == 0: + with torch.no_grad(): + _aux['w_rms'] = float(sum(p.float().pow(2).mean().sqrt() for p in params) / len(params)) + try: wb.log({'train_ce': loss.item(), 'val_ce': val, 'best': best, **_aux}, step=step) except Exception: pass if step % args.save_every == 0 or step == args.steps: torch.save({'tok': tok.state_dict(), 'pos': pos.state_dict(), 'blocks': blocks.state_dict(), |
