import sys, json, random from pathlib import Path import numpy as np import torch import torch.nn as nn import torch.nn.functional as F sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report SEEDS = tuple(range(8)) # Same union is evaluated for both systems; baseline central knob is lr and wd. GRID = [ {'lr': 0.001, 'weight_decay': 0.0}, {'lr': 0.003, 'weight_decay': 0.0}, {'lr': 0.006, 'weight_decay': 0.0}, ] class PSDConv(nn.Module): """Circular B*x followed by the exact channel-transposed flipped B* u.""" def __init__(self, n=32, m=32, k=3): super().__init__() self.B = nn.Parameter(0.08 * torch.randn(m, n, k, k)) def forward(self, x): u = F.conv2d(F.pad(x, (1,1,1,1), mode='circular'), self.B) # For cross-correlation conv, adjoint is circular correlation with spatial flip. bt = self.B.flip(-1, -2).transpose(0, 1) return F.conv2d(F.pad(u, (1,1,1,1), mode='circular'), bt) class FreeConv(nn.Module): def __init__(self, n=32, k=3): super().__init__() self.K = nn.Parameter(0.08 * torch.randn(n, n, k, k)) def forward(self, x): return F.conv2d(F.pad(x, (1,1,1,1), mode='circular'), self.K) def cnn_small_psd(out_dim, idea): # Same outer architecture as bench cnn_small; only conv2 differs. c2 = PSDConv(32, 32) if idea else FreeConv(32) return nn.Sequential( nn.Conv2d(3, 32, 3, padding=1), nn.ReLU(), nn.MaxPool2d(2), c2, nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(64, 64, 3, padding=1), nn.ReLU(), nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(64, out_dim)) # Correct channel count after c2: wrap PSD/free output into 64 channels using two banks. class BlockModel(nn.Module): def __init__(self, idea, out_dim=10): super().__init__() self.c1=nn.Conv2d(3,32,3,padding=1) self.act=nn.ReLU(); self.pool=nn.MaxPool2d(2) self.block=PSDConv(32,32) if idea else FreeConv(32) # B*B returns n=32 channels self.c3=nn.Conv2d(32,64,3,padding=1) self.head=nn.Sequential(nn.ReLU(),nn.MaxPool2d(2),nn.AdaptiveAvgPool2d(1),nn.Flatten(),nn.Linear(64,out_dim)) def forward(self,x): return self.head(self.c3(self.block(self.pool(self.act(self.c1(x))))) ) def seed_all(s): random.seed(s); np.random.seed(s); torch.manual_seed(s); torch.cuda.manual_seed_all(s) def run(idea, cfg, seed, keep=False): seed_all(seed) ds=get_dataset('vision', seed=seed, n_train=400, n_test=400) model=BlockModel(idea, ds['out_dim']) net, metric, hist=train_model(model, ds, epochs=12, lr=cfg['lr'], batch=128, weight_decay=cfg['weight_decay'], log=lambda *_: None) if net is None: return float('nan') if keep: torch.save(net.cpu().state_dict(), f'model_{"idea" if idea else "base"}_{seed}.pt') return metric def make_fn(idea, cfg): return lambda s: run(idea, cfg, int(s)) def spectrum_signature(seed, cfg): seed_all(seed); ds=get_dataset('vision',seed=seed,n_train=400,n_test=400) 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) b=net.block.B.detach().cpu().numpy() # trained-model response on a 16x16 grid; predicted PSD min eigenvalue vs observed. vals=[] for p in range(16): for q in range(16): phase=np.exp(2j*np.pi*(np.arange(3)[:,None]*p/16+np.arange(3)[None,:]*q/16)) z=(b*phase[None,None]).sum((2,3)); vals.append(np.linalg.eigvalsh(z.conj().T@z).min()) return {'predicted_min_eigenvalue':0.0,'observed_min_eigenvalue':float(min(vals)), 'predicted_nonnegative':True,'observed_nonnegative':bool(min(vals)>=-1e-5), 'confirmed':bool(min(vals)>=-1e-5)} def main(): # baseline sweep uses the same three configurations and four seed tuning budget. base=sweep_baseline(lambda c: make_fn(False,c), GRID, seeds=(0,1,2,3)) idea_trials=[] for cfg in GRID: r=evaluate(make_fn(True,cfg), SEEDS) idea_trials.append({'cfg':cfg,'result':r}) best=min(idea_trials,key=lambda z:z['result']['mean']) sig=spectrum_signature(0,best['cfg']) rep=make_report('vision','cnn_small_psd_matched',base,best['result'],{ 'idea_sweep':idea_trials, 'track_rationale':'The idea changes a spatial multichannel convolution, so CIFAR-10 vision/cnn_small is structurally matched.', 'mechanism_signature':sig}) Path('bench_report.json').write_text(json.dumps(rep,indent=2)) print(json.dumps(rep,indent=2)) if __name__=='__main__': main()