import sys, json, random from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, train_model, sweep_baseline, make_report, evaluate SEEDS = tuple(range(8)) EPOCHS = 8 BATCH = 128 class SpatialTreePool(nn.Module): def __init__(self, mode='max', p=0.5): super().__init__() self.mode, self.p = mode, float(p) self.observed_zero = [] def forward(self, x): self.observed_zero = [] a, b = x[:, :, 0::2, 0::2], x[:, :, 0::2, 1::2] c, d = x[:, :, 1::2, 0::2], x[:, :, 1::2, 1::2] def merge(u, v): if self.mode == 'max': out = torch.maximum(u, v) else: gate = torch.rand((u.shape[0], u.shape[1], 1, 1), device=u.device) < self.p out = torch.where(gate, u + v, torch.minimum(u, v)) self.observed_zero.append(float((out == 0).float().mean().detach().cpu())) return out if self.mode == 'max': # Standard CNN max pooling, retained as the baseline operation. return torch.maximum(torch.maximum(a, b), torch.maximum(c, d)) u, v = merge(a, b), merge(c, d) return merge(u, v) class TreeCNN(nn.Module): def __init__(self, out_dim, mode='max', p=.5): super().__init__() self.c1 = nn.Conv2d(3, 32, 3, padding=1) self.c2 = nn.Conv2d(32, 64, 3, padding=1) self.c3 = nn.Conv2d(64, 96, 3, padding=1) self.pool1 = SpatialTreePool(mode, p) self.pool2 = SpatialTreePool(mode, p) self.pool3 = SpatialTreePool(mode, p) self.fc1 = nn.Linear(96 * 4 * 4, 128) self.fc2 = nn.Linear(128, out_dim) self.last_zeros = [] def forward(self, x): x = torch.relu(self.c1(x)); x = self.pool1(x) x = torch.relu(self.c2(x)); x = self.pool2(x) x = torch.relu(self.c3(x)); x = self.pool3(x) self.last_zeros = self.pool1.observed_zero + self.pool2.observed_zero + self.pool3.observed_zero return self.fc2(torch.relu(self.fc1(x.flatten(1)))) def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def run_one(seed, cfg, mode, collect=False): seed_all(seed) ds = get_dataset('vision', seed, n_train=400, n_test=200) net = TreeCNN(ds['out_dim'], mode=mode, p=cfg.get('p', .5)) net, metric, hist = train_model(net, ds, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH, weight_decay=cfg.get('wd', 0.0), log=lambda *_: None) if collect: # Signature collection is diagnostic only; use CPU to avoid a second # cuDNN engine-selection failure after train_model's fallback. net = net.to('cpu'); net.eval() with torch.no_grad(): net(ds['xte'].cpu()) return float(metric), list(net.last_zeros) return float(metric) def main(): # Union parity: all lr values tried by the idea are also in baseline sweep. lrs = [1e-3, 3e-3, 1e-2] grid = [{'lr': lr, 'wd': wd} for lr in lrs for wd in [0.0, 1e-4]] base = sweep_baseline(lambda cfg: lambda s: run_one(s, cfg, 'max'), grid, seeds=(0,1,2,3)) base_full = base['full'] idea_cfgs = [{'lr': base['best_cfg']['lr'], 'wd': base['best_cfg']['wd'], 'p': p} for p in [.25,.5,.75]] idea_runs = [] for cfg in idea_cfgs: r = evaluate(lambda s, c=cfg: run_one(s, c, 'stochastic'), SEEDS) idea_runs.append((r, cfg)) idea, best_cfg = min(idea_runs, key=lambda z: z[0]['mean']) # Re-test mechanism on trained models at NN scale, using positive post-ReLU # activations and the actual stochastic tree outputs. sig = {'p': best_cfg['p'], 'q_leaf_observed': [], 'predicted_zero': [], 'observed_zero': [], 'confirmed': False} for s in SEEDS: metric, zs = run_one(s, best_cfg, 'stochastic', collect=True) if len(zs) >= 3: q = zs[0] pred = q for _ in range(2): pred = 2*(1-best_cfg['p'])*pred + (2*best_cfg['p']-1)*pred*pred sig['q_leaf_observed'].append(q); sig['predicted_zero'].append(pred); sig['observed_zero'].append(zs[-1]) if sig['observed_zero']: err = float(np.mean(np.abs(np.asarray(sig['predicted_zero']) - np.asarray(sig['observed_zero'])))) sig['mean_abs_prediction_error'] = err sig['confirmed'] = bool(err < .10) report = make_report('vision', 'cnn_small', base, idea, {'best_cfg': best_cfg, **sig}) report['idea_sweep'] = [{'cfg': c, 'mean': r['mean'], 'per_seed': r['per_seed']} for r,c in idea_runs] report['structural_match'] = 'vision CNN spatial hierarchical pooling' report['epochs'] = EPOCHS; report['n_train'] = 400; report['n_test'] = 200 Path('bench_report.json').write_text(json.dumps(report, indent=2)) print(json.dumps(report, indent=2)) if __name__ == '__main__': main()