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.py9
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(),