Doubly-Stochastic Hyper-Residual Blocks / experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, math, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6
  7SEED = 2230
  8np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
  9torch.set_num_threads(min(8, torch.get_num_threads()))
 10
 11
 12def sinkhorn_np(G, tau=1.0, steps=10):
 13    P = np.exp(np.clip(G / tau, -60, 60))
 14    for _ in range(steps):
 15        P = P / P.sum(axis=1, keepdims=True)
 16        P = P / P.sum(axis=0, keepdims=True)
 17    return P
 18
 19
 20def sinkhorn_torch(G, tau=1.0, steps=8):
 21    P = torch.exp(torch.clamp(G / tau, -30, 30))
 22    for _ in range(steps):
 23        P = P / P.sum(dim=-1, keepdim=True)
 24        P = P / P.sum(dim=-2, keepdim=True)
 25    return P
 26
 27
 28def toy_verification():
 29    rng = np.random.default_rng(SEED)
 30    S, d = 6, 13
 31    # Prediction 1: alternating normalization improves both constraint residuals.
 32    G = rng.normal(size=(S, S)) * 2.5
 33    ks = [0, 1, 2, 3, 5, 10, 20]
 34    constraint_rows = []
 35    for k in ks:
 36        A = sinkhorn_np(G, tau=0.7, steps=k)
 37        constraint_rows.append({"K": k,
 38            "max_row_error": float(np.max(np.abs(A.sum(1)-1))),
 39            "max_col_error": float(np.max(np.abs(A.sum(0)-1))),
 40            "min_entry": float(A.min())})
 41    # Prediction 2: DS mixing is non-expansive in Frobenius norm.
 42    ratios = []
 43    for scale in [0.0, 0.5, 1.0, 2.0, 5.0, 10.0]:
 44        for _ in range(40):
 45            A = sinkhorn_np(rng.normal(size=(S,S))*scale, tau=1.0, steps=100)
 46            X = rng.normal(size=(S,d))
 47            ratios.append((scale, float(np.linalg.norm(A@X)/np.linalg.norm(X))))
 48    ratio_summary = {str(s): {"max": max(r for x,r in ratios if x==s),
 49                              "mean": float(np.mean([r for x,r in ratios if x==s]))}
 50                     for s in [0.0,0.5,1.0,2.0,5.0,10.0]}
 51    # Prediction 3: exact uniform input stream mass is preserved by row-stochastic A.
 52    mass_errors = []
 53    for _ in range(100):
 54        A = sinkhorn_np(rng.normal(size=(S,S))*4, steps=100)
 55        x = rng.normal(size=S)
 56        mass_errors.append(float(abs((A@x).sum()-x.sum())))
 57    # Spectral sweep: DS matrices should have spectral norm <= 1 (up to numerical tol).
 58    spectral = []
 59    for scale in [0, .5, 1, 2, 5, 10]:
 60        vals=[]
 61        for _ in range(40):
 62            A=sinkhorn_np(rng.normal(size=(S,S))*scale, steps=100)
 63            vals.append(np.linalg.svd(A,compute_uv=False)[0])
 64        spectral.append({"scale":scale,"max_sigma":float(max(vals)),"mean_sigma":float(np.mean(vals))})
 65    return {"sinkhorn_constraints": constraint_rows, "frobenius_ratio": ratio_summary,
 66            "arbitrary_total_mass_max_error": max(mass_errors), "spectral_norm": spectral}
 67
 68
 69class StreamNet(nn.Module):
 70    def __init__(self, constrained, S=4, D=8, hidden=32, tau=1.0):
 71        super().__init__(); self.S=S; self.D=D; self.constrained=constrained; self.tau=tau
 72        self.inp=nn.Linear(1,S*D); self.g=nn.Parameter(torch.zeros(S,S))
 73        self.mlp=nn.Sequential(nn.LayerNorm(D),nn.Linear(D,hidden),nn.Tanh(),nn.Linear(hidden,D))
 74        self.out=nn.Linear(S*D,1)
 75    def forward(self,x):
 76        X=self.inp(x).view(-1,self.S,self.D)
 77        if self.constrained: A=sinkhorn_torch(self.g,self.tau,10)
 78        else: A=self.g # deliberately unconstrained control, initialized near zero
 79        # normalize control's scale only to make optimization numerically comparable
 80        if not self.constrained: A=A/(A.abs().sum(dim=1,keepdim=True)+1e-6)
 81        Xm=torch.einsum('ij,bjd->bid',A,X)
 82        F=self.mlp(X)
 83        Y=Xm + 0.1*F
 84        return self.out(Y.reshape(x.shape[0],-1)), A, X, Xm
 85
 86
 87def train_compare():
 88    torch.manual_seed(SEED)
 89    n=512; x=torch.linspace(-1,1,n).unsqueeze(1); y=torch.sin(5*x)+0.15*x*x
 90    results={}
 91    for constrained in [False,True]:
 92        torch.manual_seed(SEED)
 93        model=StreamNet(constrained)
 94        opt=torch.optim.Adam(model.parameters(),lr=3e-3)
 95        losses=[]; grad_spikes=0; mix_ratios=[]; cvs=[]
 96        for step in range(250):
 97            ix=torch.randperm(n)[:64]; pred,A,X,Xm=model(x[ix]); loss=((pred-y[ix])**2).mean()
 98            opt.zero_grad(); loss.backward()
 99            gn=float(torch.nn.utils.clip_grad_norm_(model.parameters(),1e9))
100            if gn>10: grad_spikes+=1
101            opt.step(); losses.append(float(loss.detach()))
102            with torch.no_grad():
103                mix_ratios.append(float(torch.linalg.norm(Xm)/torch.linalg.norm(X)))
104                rms=Xm.pow(2).mean(dim=(0,2)).sqrt(); cvs.append(float(rms.std()/(rms.mean()+1e-8)))
105        results['idea' if constrained else 'baseline']={
106            'final_loss':float(np.mean(losses[-25:])), 'best_loss':float(min(losses)),
107            'grad_spikes_gt10':grad_spikes, 'mix_ratio_mean':float(np.mean(mix_ratios)),
108            'stream_rms_cv_mean':float(np.mean(cvs)),
109            'row_error':float((A.sum(1)-1).abs().max()), 'col_error':float((A.sum(0)-1).abs().max())}
110    return results
111
112
113def main():
114    out={'seed':SEED,'verification':toy_verification(),'training':train_compare()}
115    Path('results.json').write_text(json.dumps(out,indent=2))
116    print(json.dumps(out,indent=2))
117
118if __name__=='__main__': main()