import sys, json, random, numpy as np, torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, train_model, make_report, sweep_baseline, evaluate SEEDS = tuple(range(8)) GRID = [ {'lr': 1e-3, 'weight_decay': 0.0}, {'lr': 3e-3, 'weight_decay': 0.0}, {'lr': 1e-3, 'weight_decay': 1e-4}, ] EPOCHS = 6 NTRAIN, NTEST = 400, 200 class GlobalSpatialNorm(nn.Module): def __init__(self, channels, eps=1e-5): super().__init__() self.gamma = nn.Parameter(torch.ones(channels)) self.beta = nn.Parameter(torch.zeros(channels)) self.eps = eps def forward(self, x): mu = x.mean(dim=(2, 3), keepdim=True) var = (x - mu).square().mean(dim=(2, 3), keepdim=True) return self.gamma[None,:,None,None] * (x-mu) / torch.sqrt(var+self.eps) + self.beta[None,:,None,None] class MatchedCNN(nn.Module): def __init__(self, out_dim=10, kind='batch'): super().__init__() norm = nn.BatchNorm2d(32) if kind == 'batch' else GlobalSpatialNorm(32) self.norm = norm self.net = nn.Sequential( nn.Conv2d(3, 32, 3, padding=1), nn.ReLU(), norm, nn.MaxPool2d(2), nn.Conv2d(32, 64, 3, padding=1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(64, 96, 3, padding=1), nn.ReLU(), nn.MaxPool2d(2), nn.Flatten(), nn.Linear(96*4*4, 128), nn.ReLU(), nn.Linear(128, out_dim)) def forward(self, x): return self.net(x) def seed_all(s): random.seed(s); np.random.seed(s); torch.manual_seed(s) def run(kind, cfg, seed, retain=False): seed_all(10000 + seed) d = get_dataset('vision', seed, n_train=NTRAIN, n_test=NTEST) model = MatchedCNN(d['out_dim'], kind) net, metric, hist = train_model(model, d, epochs=EPOCHS, lr=cfg['lr'], weight_decay=cfg['weight_decay'], batch=128, log=lambda *_: None) if net is None: raise RuntimeError('training failed') return float(metric), net, d def metric_fn(kind, cfg): return lambda seed: run(kind, cfg, seed)[0] def mechanism_signature(): seed_all(4242) d = get_dataset('vision', 0, n_train=32, n_test=8) model = MatchedCNN(d['out_dim'], 'global') model.eval() x = d['xte'][:1].clone().requires_grad_(True) # Feature immediately after the local convolution and ReLU, before classifier. h = model.net[0](x); h = model.net[1](h) t, s, ch = 0, 15, 3 scalar = model.norm(h)[0, ch, t//32, t%32] jac = torch.autograd.grad(scalar, h, retain_graph=True)[0][0, ch, s//32, s%32].item() flat = h.detach()[0, ch].reshape(-1) mu = flat.mean() sig = torch.sqrt(((flat-mu)**2).mean()+model.norm.eps) hat_t, hat_s = (flat[t]-mu)/sig, (flat[s]-mu)/sig gamma = model.norm.gamma[ch].detach() pred = (gamma/sig * (-(1+hat_t*hat_s)/flat.numel())).item() return {'location': 'trained global-normalization layer on bench vision model', 'predicted_offdiagonal': float(pred), 'observed_offdiagonal': float(jac), 'absolute_error': float(abs(pred-jac)), 'n': int(flat.numel()), 'confirmed': bool(abs(pred-jac) < 1e-5)} def main(): # Baseline sweep and idea sweep use identical configs and seeds. base = sweep_baseline(lambda cfg: metric_fn('batch', cfg), GRID) idea_trials = [] for cfg in GRID: r = evaluate(metric_fn('global', cfg), SEEDS) idea_trials.append({'cfg': cfg, 'result': r}) best = min(idea_trials, key=lambda q: q['result']['mean']) report = make_report('vision', 'cnn_small', base, best['result'], {'signature': mechanism_signature(), 'idea_sweep': idea_trials, 'matched_architecture': True, 'normalization': 'per-example channelwise spatial global statistics'}) report['custom_track'] = None with open('bench_report.json', 'w') as f: json.dump(report, f, indent=2) print(json.dumps(report, indent=2)) if __name__ == '__main__': main()