import os, sys, json, math, random import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report SEEDS = tuple(range(8)) NTR, NTE, EPOCHS, BATCH = 1000, 400, 15, 128 def seed_all(s): random.seed(s); np.random.seed(s); torch.manual_seed(s) if torch.cuda.is_available(): torch.cuda.manual_seed_all(s) def device(): return 'cuda' if torch.cuda.is_available() else 'cpu' def smooth_polar_torch(m, eps): # thin SVD; equivalent to U diag(s/sqrt(s^2+eps)) V^T u, s, vh = torch.linalg.svd(m, full_matrices=False) return (u * (s / torch.sqrt(s*s + eps))) @ vh def train(seed, cfg, method, collect=False): seed_all(seed) ds = get_dataset('tabular', seed, n_train=NTR, n_test=NTE) net = make_model('mlp_tiny', ds['input_shape'], ds['out_dim']) dev = device() try: net = net.to(dev) x, y = ds['xtr'].to(dev), ds['ytr'].to(dev) lossf = nn.MSELoss() if method == 'adam': opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'], betas=(cfg['beta1'], 0.999), weight_decay=cfg.get('weight_decay', 0.0)) mom = None else: opt = None mom = {id(p): torch.zeros_like(p) for p in net.parameters() if p.ndim == 2} response_ratios, response_observed, update_norms = [], [], [] for ep in range(EPOCHS): net.train(); perm = torch.randperm(len(x), device=dev) for start in range(0, len(x), BATCH): idx = perm[start:start+BATCH] loss = lossf(net(x[idx]), y[idx]) net.zero_grad(set_to_none=True); loss.backward() if method == 'adam': opt.step() else: total_update = 0.0 for p in net.parameters(): if p.grad is None: continue if p.ndim == 2: z = mom[id(p)] z.mul_(cfg['beta']).add_(p.grad) eps = cfg['c'] * torch.mean(z*z).detach().clamp_min(1e-20) upd = smooth_polar_torch(z, eps) p.data.add_(upd, alpha=-cfg['lr']) if collect and ep == EPOCHS-1: s = torch.linalg.svdvals(z).detach() # Stage-1 prediction r(t)=t/sqrt(t^2+1), tested at t=1. t = s / torch.sqrt(eps) pred = t / torch.sqrt(t*t + 1) obs = torch.linalg.svdvals(upd).detach() response_ratios.extend(t.cpu().numpy().tolist()) response_observed.extend(obs.cpu().numpy().tolist()) total_update += float(torch.linalg.norm(upd).detach())**2 else: # Non-matrix parameters remain ordinary SGD, as specified. p.data.add_(p.grad, alpha=-cfg['lr']) if collect: update_norms.append(math.sqrt(total_update)) net.eval() with torch.no_grad(): metric = float(torch.mean((net(ds['xte'].to(dev)) - ds['yte'].to(dev))**2)) stats = {} if collect and response_ratios: rr, oo = np.asarray(response_ratios), np.asarray(response_observed) pred = rr / np.sqrt(rr*rr + 1.0) stats = {'predicted_response_mean': float(pred.mean()), 'observed_response_mean': float(oo.mean()), 'response_mae': float(np.mean(np.abs(pred-oo))), 'update_norm_cv': float(np.std(update_norms)/(np.mean(update_norms)+1e-12)), 'n_observations': int(len(rr))} return metric, stats except (RuntimeError, torch.cuda.OutOfMemoryError): # Robust CPU fallback with the same seeded model/data. if dev == 'cuda': torch.cuda.empty_cache() return train_cpu(seed, cfg, method, collect) raise def train_cpu(seed, cfg, method, collect=False): old = torch.cuda.is_available # A direct CPU implementation avoids relying on CUDA state after an error. seed_all(seed); ds=get_dataset('tabular',seed,n_train=NTR,n_test=NTE) net=make_model('mlp_tiny',ds['input_shape'],ds['out_dim']); x,y=ds['xtr'],ds['ytr']; mom={id(p):torch.zeros_like(p) for p in net.parameters() if p.ndim==2}; opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'],betas=(cfg.get('beta1',.9),.999)) if method=='adam' else None for _ in range(EPOCHS): for st in range(0,len(x),BATCH): net.zero_grad(); loss=nn.functional.mse_loss(net(x[st:st+BATCH]),y[st:st+BATCH]); loss.backward() if method=='adam': opt.step() else: for p in net.parameters(): if p.grad is None: continue if p.ndim==2: z=mom[id(p)]; z.mul_(cfg['beta']).add_(p.grad); e=cfg['c']*torch.mean(z*z).clamp_min(1e-20); p.data.add_(smooth_polar_torch(z,e),alpha=-cfg['lr']) else: p.data.add_(p.grad,alpha=-cfg['lr']) with torch.no_grad(): return float(nn.functional.mse_loss(net(ds['xte']),ds['yte'])), {} def main(): # Union of all lrs used by either side; Adam's beta1 is its central method knob. lrs=[0.0015,0.003,0.006] base_grid=[{'lr':lr,'beta1':b,'weight_decay':0.0} for lr in lrs for b in (0.85,0.95)] def base_factory(cfg): return lambda s: train(s,cfg,'adam')[0] baseline=sweep_baseline(base_factory,base_grid,seeds=(0,1,2,3)) idea_cfgs=[{'lr':lr,'beta':0.9,'c':0.001} for lr in lrs] idea_records=[]; best_cfg=None; best_mean=float('inf') for cfg in idea_cfgs: r=evaluate(lambda s: train(s,cfg,'smooth')[0], seeds=SEEDS) idea_records.append({'cfg':cfg,'result':r}) if r['mean']