summaryrefslogtreecommitdiff
path: root/ep_run/cascade_probe.py
diff options
context:
space:
mode:
Diffstat (limited to 'ep_run/cascade_probe.py')
-rw-r--r--ep_run/cascade_probe.py22
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')