Critical stochastic min-plus tree layer / bench_experiment.py

Failed on benchmark

Raw ⬇ ZIP
  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()