PSD Spectral CNN Block / stage2_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, json, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6import torch.nn.functional as F
  7
  8sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  9from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
 10
 11SEEDS = tuple(range(8))
 12# Same union is evaluated for both systems; baseline central knob is lr and wd.
 13GRID = [
 14    {'lr': 0.001, 'weight_decay': 0.0},
 15    {'lr': 0.003, 'weight_decay': 0.0},
 16    {'lr': 0.006, 'weight_decay': 0.0},
 17]
 18
 19class PSDConv(nn.Module):
 20    """Circular B*x followed by the exact channel-transposed flipped B* u."""
 21    def __init__(self, n=32, m=32, k=3):
 22        super().__init__()
 23        self.B = nn.Parameter(0.08 * torch.randn(m, n, k, k))
 24    def forward(self, x):
 25        u = F.conv2d(F.pad(x, (1,1,1,1), mode='circular'), self.B)
 26        # For cross-correlation conv, adjoint is circular correlation with spatial flip.
 27        bt = self.B.flip(-1, -2).transpose(0, 1)
 28        return F.conv2d(F.pad(u, (1,1,1,1), mode='circular'), bt)
 29
 30class FreeConv(nn.Module):
 31    def __init__(self, n=32, k=3):
 32        super().__init__()
 33        self.K = nn.Parameter(0.08 * torch.randn(n, n, k, k))
 34    def forward(self, x):
 35        return F.conv2d(F.pad(x, (1,1,1,1), mode='circular'), self.K)
 36
 37def cnn_small_psd(out_dim, idea):
 38    # Same outer architecture as bench cnn_small; only conv2 differs.
 39    c2 = PSDConv(32, 32) if idea else FreeConv(32)
 40    return nn.Sequential(
 41        nn.Conv2d(3, 32, 3, padding=1), nn.ReLU(), nn.MaxPool2d(2),
 42        c2, nn.ReLU(), nn.MaxPool2d(2),
 43        nn.Conv2d(64, 64, 3, padding=1), nn.ReLU(), nn.AdaptiveAvgPool2d(1),
 44        nn.Flatten(), nn.Linear(64, out_dim))
 45
 46# Correct channel count after c2: wrap PSD/free output into 64 channels using two banks.
 47class BlockModel(nn.Module):
 48    def __init__(self, idea, out_dim=10):
 49        super().__init__()
 50        self.c1=nn.Conv2d(3,32,3,padding=1)
 51        self.act=nn.ReLU(); self.pool=nn.MaxPool2d(2)
 52        self.block=PSDConv(32,32) if idea else FreeConv(32) # B*B returns n=32 channels
 53        self.c3=nn.Conv2d(32,64,3,padding=1)
 54        self.head=nn.Sequential(nn.ReLU(),nn.MaxPool2d(2),nn.AdaptiveAvgPool2d(1),nn.Flatten(),nn.Linear(64,out_dim))
 55    def forward(self,x): return self.head(self.c3(self.block(self.pool(self.act(self.c1(x))))) )
 56
 57def seed_all(s):
 58    random.seed(s); np.random.seed(s); torch.manual_seed(s); torch.cuda.manual_seed_all(s)
 59
 60def run(idea, cfg, seed, keep=False):
 61    seed_all(seed)
 62    ds=get_dataset('vision', seed=seed, n_train=400, n_test=400)
 63    model=BlockModel(idea, ds['out_dim'])
 64    net, metric, hist=train_model(model, ds, epochs=12, lr=cfg['lr'], batch=128,
 65                                  weight_decay=cfg['weight_decay'], log=lambda *_: None)
 66    if net is None: return float('nan')
 67    if keep:
 68        torch.save(net.cpu().state_dict(), f'model_{"idea" if idea else "base"}_{seed}.pt')
 69    return metric
 70
 71def make_fn(idea, cfg): return lambda s: run(idea, cfg, int(s))
 72
 73def spectrum_signature(seed, cfg):
 74    seed_all(seed); ds=get_dataset('vision',seed=seed,n_train=400,n_test=400)
 75    model=BlockModel(True,ds['out_dim']); net,_,_=train_model(model,ds,epochs=12,lr=cfg['lr'],batch=128,weight_decay=cfg['weight_decay'],log=lambda *_:None)
 76    b=net.block.B.detach().cpu().numpy()
 77    # trained-model response on a 16x16 grid; predicted PSD min eigenvalue vs observed.
 78    vals=[]
 79    for p in range(16):
 80      for q in range(16):
 81        phase=np.exp(2j*np.pi*(np.arange(3)[:,None]*p/16+np.arange(3)[None,:]*q/16))
 82        z=(b*phase[None,None]).sum((2,3)); vals.append(np.linalg.eigvalsh(z.conj().T@z).min())
 83    return {'predicted_min_eigenvalue':0.0,'observed_min_eigenvalue':float(min(vals)),
 84            'predicted_nonnegative':True,'observed_nonnegative':bool(min(vals)>=-1e-5),
 85            'confirmed':bool(min(vals)>=-1e-5)}
 86
 87def main():
 88    # baseline sweep uses the same three configurations and four seed tuning budget.
 89    base=sweep_baseline(lambda c: make_fn(False,c), GRID, seeds=(0,1,2,3))
 90    idea_trials=[]
 91    for cfg in GRID:
 92        r=evaluate(make_fn(True,cfg), SEEDS)
 93        idea_trials.append({'cfg':cfg,'result':r})
 94    best=min(idea_trials,key=lambda z:z['result']['mean'])
 95    sig=spectrum_signature(0,best['cfg'])
 96    rep=make_report('vision','cnn_small_psd_matched',base,best['result'],{
 97      'idea_sweep':idea_trials,
 98      'track_rationale':'The idea changes a spatial multichannel convolution, so CIFAR-10 vision/cnn_small is structurally matched.',
 99      'mechanism_signature':sig})
100    Path('bench_report.json').write_text(json.dumps(rep,indent=2))
101    print(json.dumps(rep,indent=2))
102if __name__=='__main__': main()