diff options
Diffstat (limited to 'ep_run/cascade_probe.py')
| -rw-r--r-- | ep_run/cascade_probe.py | 22 |
1 files changed, 19 insertions, 3 deletions
diff --git a/ep_run/cascade_probe.py b/ep_run/cascade_probe.py index d6d9fea..a6bb3aa 100644 --- a/ep_run/cascade_probe.py +++ b/ep_run/cascade_probe.py @@ -26,11 +26,18 @@ ap.add_argument('--ckpt', type=str, default='') # gate at a casc_bp_trai ap.add_argument('--sumscale', action='store_true') # relax in SUM units (gamma=1 Z-IL semantics) ap.add_argument('--init_sweep', action='store_true') # one reverse gamma=1 sweep as state INIT (readout still at equilibrium = clean EP) ap.add_argument('--sopt', choices=['sgd', 'adam'], default='sgd') # state-relaxation optimizer +ap.add_argument('--prec', choices=['fp32', 'tf32', 'bf16'], default='fp32') # A0.4 precision gate (EP side only; BP ref stays fp32) ap.add_argument('--geta_auto', type=float, default=0.0) # >0: per-layer gamma_l = c/(1+sigma_l+1^2), c=this; sigma via power-iter args = ap.parse_args() torch.manual_seed(args.seed) dev = 'cuda' if torch.cuda.is_available() else 'cpu' -torch.set_float32_matmul_precision('highest') +if args.prec == 'fp32': + torch.set_float32_matmul_precision('highest') +elif args.prec == 'tf32': + torch.backends.cuda.matmul.allow_tf32 = True + torch.backends.cudnn.allow_tf32 = True + torch.set_float32_matmul_precision('high') +_BF16 = (args.prec == 'bf16') DD = Path('/home/yurenh2/ept/ep_run/data/tinystories_bpe') vocab = pickle.load(open(DD / 'meta.pkl', 'rb'))['vocab_size'] @@ -237,7 +244,11 @@ for bi in range(args.batches): ce = F.cross_entropy(readout(zs[-1]).reshape(-1, vocab), y.reshape(-1)) gbp = list(torch.autograd.grad(ce, gate_params, allow_unused=True)) gbp = [g if g is not None else torch.zeros(1, device=dev) for g in gbp] - # two-phase (or single-phase interleaved zil) + # EP side: pure-bf16 MODEL when --prec bf16 (BP reference above stays fp32); cast back after + if _BF16: + sd32 = [ (m, {k: v.clone() for k, v in m.state_dict().items()}) for m in [tok, pos, blocks] ] + for m in [tok, pos, blocks]: m.bfloat16() + z0 = z0.bfloat16(); mask = mask.bfloat16() if args.scheme == 'zil': gep = zil_grads(x, y, args.beta) else: @@ -245,6 +256,11 @@ for bi in range(args.batches): gp = dFdtheta(zp, z0, y, +args.beta, gate_params, x) gm = dFdtheta(zm, z0, y, -args.beta, gate_params, x) gep = [(a - b) / (2 * args.beta) for a, b in zip(gp, gm)] + gep = [g.float() for g in gep] + if _BF16: + for m, sd in sd32: + m.float(); m.load_state_dict(sd) + mask = mask.float() cos_b.append(F.cosine_similarity(flat(gep), flat(gbp), dim=0).item()) shr_b.append((flat(gep).norm() / flat(gbp).norm()).item()) for l in range(args.L): @@ -257,6 +273,6 @@ print('per-block cos: ' + ' '.join(f'{l}:{np.mean(v):.4f}' for l, v in enumerat if args.include_io: idx = [i for i, n in enumerate(names) if n.startswith('io.')] print(f'io cos (last batch): {F.cosine_similarity(flat([gep[i] for i in idx]), flat([gbp[i] for i in idx]), dim=0).item():.4f}') -print(f'SUMMARY scheme={args.scheme}{"+init" if args.init_sweep else ""}+{args.sopt} L={args.L} C={args.C} K={args.K} eta={args.eta} beta={args.beta} ' +print(f'SUMMARY prec={args.prec} scheme={args.scheme}{"+init" if args.init_sweep else ""}+{args.sopt} L={args.L} C={args.C} K={args.K} eta={args.eta} beta={args.beta} ' f'io={int(args.include_io)} warp={args.warp} ckpt={args.ckpt or "-"} ' f'cos={cm:.4f} cosmin={cmin:.4f} shrink={sm:.3f} t={time.time()-t0:.1f}s') |
