Critical stochastic min-plus tree layer / bench_experiment.py
Failed on benchmark
1import sys, json, random
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6
7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
8from bench import get_dataset, train_model, sweep_baseline, make_report, evaluate
9
10SEEDS = tuple(range(8))
11EPOCHS = 8
12BATCH = 128
13
14class SpatialTreePool(nn.Module):
15 def __init__(self, mode='max', p=0.5):
16 super().__init__()
17 self.mode, self.p = mode, float(p)
18 self.observed_zero = []
19
20 def forward(self, x):
21 self.observed_zero = []
22 a, b = x[:, :, 0::2, 0::2], x[:, :, 0::2, 1::2]
23 c, d = x[:, :, 1::2, 0::2], x[:, :, 1::2, 1::2]
24 def merge(u, v):
25 if self.mode == 'max':
26 out = torch.maximum(u, v)
27 else:
28 gate = torch.rand((u.shape[0], u.shape[1], 1, 1), device=u.device) < self.p
29 out = torch.where(gate, u + v, torch.minimum(u, v))
30 self.observed_zero.append(float((out == 0).float().mean().detach().cpu()))
31 return out
32 if self.mode == 'max':
33 # Standard CNN max pooling, retained as the baseline operation.
34 return torch.maximum(torch.maximum(a, b), torch.maximum(c, d))
35 u, v = merge(a, b), merge(c, d)
36 return merge(u, v)
37
38class TreeCNN(nn.Module):
39 def __init__(self, out_dim, mode='max', p=.5):
40 super().__init__()
41 self.c1 = nn.Conv2d(3, 32, 3, padding=1)
42 self.c2 = nn.Conv2d(32, 64, 3, padding=1)
43 self.c3 = nn.Conv2d(64, 96, 3, padding=1)
44 self.pool1 = SpatialTreePool(mode, p)
45 self.pool2 = SpatialTreePool(mode, p)
46 self.pool3 = SpatialTreePool(mode, p)
47 self.fc1 = nn.Linear(96 * 4 * 4, 128)
48 self.fc2 = nn.Linear(128, out_dim)
49 self.last_zeros = []
50
51 def forward(self, x):
52 x = torch.relu(self.c1(x)); x = self.pool1(x)
53 x = torch.relu(self.c2(x)); x = self.pool2(x)
54 x = torch.relu(self.c3(x)); x = self.pool3(x)
55 self.last_zeros = self.pool1.observed_zero + self.pool2.observed_zero + self.pool3.observed_zero
56 return self.fc2(torch.relu(self.fc1(x.flatten(1))))
57
58def seed_all(seed):
59 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
60 if torch.cuda.is_available():
61 torch.cuda.manual_seed_all(seed)
62
63def run_one(seed, cfg, mode, collect=False):
64 seed_all(seed)
65 ds = get_dataset('vision', seed, n_train=400, n_test=200)
66 net = TreeCNN(ds['out_dim'], mode=mode, p=cfg.get('p', .5))
67 net, metric, hist = train_model(net, ds, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH, weight_decay=cfg.get('wd', 0.0), log=lambda *_: None)
68 if collect:
69 # Signature collection is diagnostic only; use CPU to avoid a second
70 # cuDNN engine-selection failure after train_model's fallback.
71 net = net.to('cpu'); net.eval()
72 with torch.no_grad():
73 net(ds['xte'].cpu())
74 return float(metric), list(net.last_zeros)
75 return float(metric)
76
77def main():
78 # Union parity: all lr values tried by the idea are also in baseline sweep.
79 lrs = [1e-3, 3e-3, 1e-2]
80 grid = [{'lr': lr, 'wd': wd} for lr in lrs for wd in [0.0, 1e-4]]
81 base = sweep_baseline(lambda cfg: lambda s: run_one(s, cfg, 'max'), grid, seeds=(0,1,2,3))
82 base_full = base['full']
83 idea_cfgs = [{'lr': base['best_cfg']['lr'], 'wd': base['best_cfg']['wd'], 'p': p} for p in [.25,.5,.75]]
84 idea_runs = []
85 for cfg in idea_cfgs:
86 r = evaluate(lambda s, c=cfg: run_one(s, c, 'stochastic'), SEEDS)
87 idea_runs.append((r, cfg))
88 idea, best_cfg = min(idea_runs, key=lambda z: z[0]['mean'])
89 # Re-test mechanism on trained models at NN scale, using positive post-ReLU
90 # activations and the actual stochastic tree outputs.
91 sig = {'p': best_cfg['p'], 'q_leaf_observed': [], 'predicted_zero': [], 'observed_zero': [], 'confirmed': False}
92 for s in SEEDS:
93 metric, zs = run_one(s, best_cfg, 'stochastic', collect=True)
94 if len(zs) >= 3:
95 q = zs[0]
96 pred = q
97 for _ in range(2): pred = 2*(1-best_cfg['p'])*pred + (2*best_cfg['p']-1)*pred*pred
98 sig['q_leaf_observed'].append(q); sig['predicted_zero'].append(pred); sig['observed_zero'].append(zs[-1])
99 if sig['observed_zero']:
100 err = float(np.mean(np.abs(np.asarray(sig['predicted_zero']) - np.asarray(sig['observed_zero']))))
101 sig['mean_abs_prediction_error'] = err
102 sig['confirmed'] = bool(err < .10)
103 report = make_report('vision', 'cnn_small', base, idea, {'best_cfg': best_cfg, **sig})
104 report['idea_sweep'] = [{'cfg': c, 'mean': r['mean'], 'per_seed': r['per_seed']} for r,c in idea_runs]
105 report['structural_match'] = 'vision CNN spatial hierarchical pooling'
106 report['epochs'] = EPOCHS; report['n_train'] = 400; report['n_test'] = 200
107 Path('bench_report.json').write_text(json.dumps(report, indent=2))
108 print(json.dumps(report, indent=2))
109
110if __name__ == '__main__': main()