Invariant-Sphere Recurrent State / run_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import json, random
 2from pathlib import Path
 3import numpy as np
 4
 5SEED = 7
 6np.random.seed(SEED); random.seed(SEED)
 7
 8# Ring cubic field from the proposal.  This coefficient choice is dissipative:
 9# <u,Q(u)> = -sum_i (u_i^4 + u_i^2 u_{i+1}^2) < 0.
10def ring_cubic(h, a=-1.0, b=0.0, c=-1.0, d=0.0):
11    z = np.roll(h, -1, axis=-1)
12    return a*h**3 + b*h**2*z + c*h*z**2 + d*z**3
13
14def radial_coeff(u, a=-1., b=0., c=-1., d=0.):
15    return np.sum(u * ring_cubic(u, a, b, c, d), axis=-1)
16
17def math_check(n=16, samples=20000):
18    u = np.random.randn(samples, n)
19    u /= np.linalg.norm(u, axis=1, keepdims=True)
20    inn = radial_coeff(u)
21    # The valid claim for n>1 is negativity, not constant -kappa.
22    # Holder/Cauchy gives sum u_i^4 >= 1/n, hence the coefficient is <= -1/n.
23    kappa_bound = 1.0/n
24    lam, dt = 1., .001
25    # Check the Lyapunov derivative directly along several Euler trajectories.
26    violations = 0
27    final_radii = []
28    initial_radii = []
29    for j in range(32):
30        h = (0.2 + 3.0*np.random.rand()) * np.random.randn(n)
31        initial_radii.append(np.linalg.norm(h))
32        for _ in range(10000):
33            h += dt*(lam*h + ring_cubic(h[None,:])[0])
34        final_radii.append(np.linalg.norm(h))
35    # For the scalar n=1 specialization, q=-y^3 is exactly the claimed logistic radial ODE.
36    h = np.array([2.0]); radii=[]; times=[]
37    for k in range(5000):
38        radii.append(abs(h[0])); times.append(k*dt)
39        h += dt*(h + (-h**3))
40    radii=np.asarray(radii); times=np.asarray(times)
41    mask=(times>.2)&(times<1.5)&(np.abs(radii-1)>1e-7)
42    slope=float(np.polyfit(times[mask], np.log(np.abs(radii[mask]-1)), 1)[0])
43    return {
44      "radial_mean":float(inn.mean()), "radial_std":float(inn.std()),
45      "radial_max":float(inn.max()), "theoretical_upper_bound":-kappa_bound,
46      "positive_radial_fraction":float(np.mean(inn>=0)),
47      "trajectory_initial_radius_range":[float(min(initial_radii)),float(max(initial_radii))],
48      "trajectory_final_radius_range":[float(min(final_radii)),float(max(final_radii))],
49      "scalar_exact_final_radius":float(radii[-1]),
50      "scalar_relaxation_log_slope":slope, "scalar_predicted_slope":-2.0,
51      "pole_multiplier_c_plus_1":{str(c):c+1 for c in [-1.5,-1.,-.5,0.]}}
52
53def train_toy():
54    try:
55        import torch
56        torch.set_num_threads(2)
57        torch.manual_seed(SEED)
58        device=torch.device("cuda" if torch.cuda.is_available() else "cpu")
59        n,T,batch,steps=8,40,32,120
60        def make_model():
61            return [torch.nn.Parameter(torch.randn(n,1,device=device)*.25),
62                    torch.nn.Parameter(torch.randn(1,n,device=device)*.25),
63                    torch.nn.Parameter(torch.zeros(n,device=device)),
64                    torch.nn.Parameter(torch.zeros(1,device=device))]
65        def forward(par,x,kind):
66            W,V,b,out=par; h=torch.zeros(x.shape[0],n,device=device); dt=.05
67            for t in range(T):
68                inp=x[:,t:t+1]@W.T
69                if kind=='tanh': h=torch.tanh(h+inp+b)
70                else:
71                    z=torch.roll(h,-1,dims=1); Q=-h**3-h*z**2
72                    h=h+dt*(h+Q+inp+b)
73            return (h@V.T).squeeze(1)+out,h
74        results={}
75        for kind in ('tanh','sphere'):
76            torch.manual_seed(SEED+(kind=='sphere')); par=make_model()
77            opt=torch.optim.Adam(par,lr=.01); losses=[]
78            for step in range(steps):
79                bit=(torch.randint(0,2,(batch,),device=device)*2-1).float()
80                x=torch.zeros(batch,T,device=device); x[:,0]=bit
81                pred,h=forward(par,x,kind); loss=((pred-bit)**2).mean()
82                opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(par,5.); opt.step()
83                if step%20==0: losses.append(float(loss.detach().cpu()))
84            bit=(torch.randint(0,2,(256,),device=device)*2-1).float(); x=torch.zeros(256,T,device=device); x[:,0]=bit
85            with torch.no_grad(): pred,h=forward(par,x,kind)
86            results[kind]={'test_mse':float(((pred-bit)**2).mean().cpu()),
87              'training_checkpoints':losses,
88              'final_hidden_norm_mean':float(h.norm(dim=1).mean().cpu()),
89              'final_hidden_norm_std':float(h.norm(dim=1).std().cpu()),'device':str(device)}
90        return results
91    except Exception as e: return {'error':repr(e)}
92
93if __name__=='__main__':
94    out={'seed':SEED,'math':math_check(),'toy':train_toy()}
95    Path('results.json').write_text(json.dumps(out,indent=2)); print(json.dumps(out,indent=2))