Two-Channel Fractal Renormalization Network / bench_run.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1import sys, json, math, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5from torch import nn
  6import torch.nn.functional as F
  7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  8from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
  9
 10SEED0 = 1723
 11LRS = [1e-3, 3e-3, 1e-2]
 12EPOCHS = 5
 13NTR, NTE = 400, 200
 14
 15class TernaryMeanPool(nn.Module):
 16    """Standard pooling baseline used by the bench CNN."""
 17    def forward(self, x):
 18        return F.max_pool2d(x, 2)
 19
 20class FractalPool(nn.Module):
 21    """Two-channel ternary renormalization pooling, with bounded defect route."""
 22    def __init__(self, channels, width=16):
 23        super().__init__()
 24        self.x_raw = nn.Parameter(torch.tensor(-1.0))
 25        self.lam_raw = nn.Parameter(torch.tensor(-0.15))
 26        self.defect = nn.Sequential(nn.Conv2d(3*channels, width, 1), nn.GELU(),
 27                                    nn.Conv2d(width, channels, 1), nn.Tanh())
 28        self.gate = nn.Sequential(nn.Conv2d(3*channels, channels, 1), nn.Sigmoid())
 29        self.last_rho = 0.0
 30    def forward(self, x):
 31        if x.shape[-1] % 2 or x.shape[-2] % 2:
 32            x = F.pad(x, (0, x.shape[-1] % 2, 0, x.shape[-2] % 2))
 33        a = x[..., 0::2, 0::2]; b = x[..., 0::2, 1::2]; c = x[..., 1::2, 0::2]
 34        # neutral and defect channels: disagreement is a child-minus-neutral signal.
 35        neutral0 = (a + b + c) / 3.0
 36        da, db, dc = a-neutral0, b-neutral0, c-neutral0
 37        lam, xx = F.softplus(self.lam_raw), F.softplus(self.x_raw)
 38        neutral = lam.pow(3)*a*b*c + 2*xx.pow(3)*da*db*dc
 39        neutral = neutral / (neutral.square().mean(1, keepdim=True).add(1e-5).sqrt())
 40        defect_in = torch.cat((da, db, dc), 1)
 41        proposed = self.defect(defect_in)
 42        gate = self.gate(defect_in)
 43        defect = gate*proposed + (1-gate)*neutral0
 44        self._last_neutral, self._last_defect = neutral, defect
 45        return neutral + defect
 46    def coeffs(self):
 47        return float(F.softplus(self.x_raw)), float(F.softplus(self.lam_raw))
 48
 49class CNN(nn.Module):
 50    def __init__(self, idea=False):
 51        super().__init__(); P = FractalPool if idea else TernaryMeanPool
 52        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)
 53        self.p1=P(32) if idea else P(); self.p2=P(64) if idea else P(); self.p3=P(96) if idea else P()
 54        self.fc=nn.Sequential(nn.Flatten(),nn.Linear(96*4*4,128),nn.ReLU(),nn.Linear(128,10))
 55    def forward(self,x):
 56        x=F.relu(self.c1(x)); x=self.p1(x); x=F.relu(self.c2(x)); x=self.p2(x); x=F.relu(self.c3(x)); x=self.p3(x); return self.fc(x)
 57
 58def math_check():
 59    a,b=0.43,0.27; lam,x=0.8,0.35
 60    f=lambda u: lam**3*u[0]**3+2*x**3*u[1]**3
 61    eps=1e-5; fd=np.array([(f((a+eps,b))-f((a-eps,b)))/(2*eps),(f((a,b+eps))-f((a,b-eps)))/(2*eps)])
 62    exact=np.array([3*lam**3*a*a,6*x**3*b*b])
 63    return {'analytic_jacobian':exact.tolist(),'finite_difference':fd.tolist(),'max_abs_error':float(np.max(np.abs(exact-fd))), 'cubic_gain_ratio':float(exact[0]/(3*a*a))}
 64
 65def run_cfg(idea, lr, seed, keep=False):
 66    torch.manual_seed(10000+seed); np.random.seed(10000+seed); random.seed(10000+seed)
 67    d=get_dataset('vision',seed,n_train=NTR,n_test=NTE); m=CNN(idea=idea)
 68    net,metric,hist=train_model(m,d,epochs=EPOCHS,lr=lr,batch=128,log=lambda *_:None)
 69    return (float(metric),net,d) if keep else float(metric)
 70
 71def full(idea,lr):
 72    vals=[]; models=[]
 73    for s in range(8):
 74        r=run_cfg(idea,lr,s,keep=idea)
 75        if idea: vals.append(r[0]); models.append(r)
 76        else: vals.append(r)
 77    return {'mean':float(np.mean(vals)),'std':float(np.std(vals)),'per_seed':vals,'n':len(vals)}, models
 78
 79def signature(models):
 80    ratios=[]; rhos=[]
 81    for metric,net,d in models:
 82        net = net.to('cpu'); net.eval(); x=d['xte'][:8].to('cpu').clone(); noise=torch.randn_like(x)*1e-3
 83        with torch.no_grad():
 84            # Re-test trained network behavior, not an analytic toy: first pooling perturbation.
 85            z=F.relu(net.c1(x)); zn=F.relu(net.c1(x+noise));
 86            y=net.p1(z); yn=net.p1(zn)
 87            ratios.append(float((yn-y).norm()/(noise.norm()+1e-12)))
 88        rhos.append(net.p1.coeffs() if hasattr(net.p1,'coeffs') else None)
 89    vals=np.asarray(ratios)
 90    return {'prediction':'local Jacobian controls perturbation amplification; cubic neutral derivative scales as lambda^3',
 91            'observed_trained_pool_perturbation_gain_mean':float(vals.mean()),
 92            'observed_trained_pool_perturbation_gain_std':float(vals.std()),
 93            'trained_coefficients':rhos, 'predicted_lambda_cubic_gain':True,
 94            'confirmed': False}
 95
 96def main():
 97    check=math_check(); print('MATH_CHECK',json.dumps(check))
 98    # Baseline sweep uses exactly the union of idea learning rates; its full result is best config.
 99    base=sweep_baseline(lambda cfg: (lambda s: run_cfg(False,cfg['lr'],s)), [{'lr':x} for x in LRS], seeds=(0,1,2,3))
100    base_full_by_lr={str(lr):full(False,lr)[0] for lr in LRS}
101    idea_candidates=[]
102    for lr in LRS:
103        r,_=full(True,lr); idea_candidates.append((r['mean'],lr,r))
104    _,best_lr,idea_res=min(idea_candidates,key=lambda z:z[0]); _,models=full(True,best_lr)
105    rep=make_report('vision','cnn_small',base,idea_res,{'math_check':check,'trained_behavior':signature(models),'baseline_full_by_lr':base_full_by_lr,'idea_sweep':[{'lr':lr,'mean':r['mean']} for _,lr,r in idea_candidates]})
106    rep['protocol_note']='Vision selected because the idea changes convolutional hierarchical feature propagation; baseline and idea share convolution/classifier stages and differ only in pooling mechanism.'
107    Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
108if __name__=='__main__': main()