Doubly-Stochastic Hyper-Residual Blocks / experiment.py
Mechanism confirmed, baseline not beaten
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()