ISS-Gated Positive Neural State Module / iss_experiment.py
Failed on benchmark
1import json
2from pathlib import Path
3import numpy as np
4
5
6def free_energy(x, x_star):
7 x = np.maximum(np.asarray(x, dtype=float), 1e-15)
8 xs = np.asarray(x_star, dtype=float)
9 v = x * np.log(x / xs) - x + xs
10 return np.sum(v, axis=-1) if x.ndim else float(v)
11
12
13class PositiveMassAction:
14 """Positive mass-action state block with RK4 integration.
15
16 reactants[j,i]=v_ij and products[j,i]=v'_ij. Rates are positive.
17 """
18 def __init__(self, reactants, products, dt=0.2, substeps=2, eps=1e-6):
19 self.reactants = np.asarray(reactants, float)
20 self.stoich = np.asarray(products, float) - self.reactants
21 self.dt, self.substeps, self.eps = dt, substeps, eps
22
23 def field(self, x, rates):
24 x = np.maximum(np.asarray(x, float), self.eps)
25 rates = np.maximum(np.asarray(rates, float), self.eps)
26 monomial = np.prod(x[..., None, :] ** self.reactants, axis=-1)
27 return np.sum(monomial[..., :, None] * rates[..., :, None] * self.stoich, axis=-2)
28
29 def step(self, x, rates):
30 x = np.asarray(x, float)
31 h = self.dt / self.substeps
32 for _ in range(self.substeps):
33 k1 = self.field(x, rates)
34 k2 = self.field(x + h*k1/2, rates)
35 k3 = self.field(x + h*k2/2, rates)
36 k4 = self.field(x + h*k3, rates)
37 x = np.maximum(x + h*(k1+2*k2+2*k3+k4)/6, self.eps)
38 return x
39
40 def iss_residual(self, x, x_next, rates, rates_star, c, k):
41 v = free_energy(x, x_star=1.0)
42 v_next = free_energy(x_next, x_star=1.0)
43 d2 = np.sum((np.asarray(rates)-np.asarray(rates_star))**2)
44 return (v_next-v)/self.dt + c*v - k*d2
45
46
47def euler(x, u, k, dt):
48 return x + dt*(u-k*x)
49
50
51def rk4_birth_death(x, u, k, dt):
52 f=lambda z: u-k*z
53 a=f(x); b=f(x+dt*a/2); c=f(x+dt*b/2); d=f(x+dt*c)
54 return x+dt*(a+2*b+2*c+d)/6
55
56
57def verify():
58 k=u0=xs=1.0
59 # Prediction 1: Euler multiplier is 1-k*dt, so stability boundary is dt*k=2.
60 lo, hi = 0., 3.
61 for _ in range(45):
62 dt=(lo+hi)/2; x=1.2
63 for _ in range(150): x=euler(x,u0,k,dt)
64 if abs(x-xs)<1e-4: lo=dt
65 else: hi=dt
66 # Prediction 2: shifted equilibrium free energy is quadratic in rate amplitude.
67 amps=np.array([.01,.02,.04,.08,.16,.25,.35])
68 Vs=np.array([free_energy((u0+a)/k,xs) for a in amps])
69 slope=np.polyfit(np.log(amps),np.log(Vs),1)[0]
70 # Prediction 3: local V decays at twice the state decay rate.
71 x=1.2; dt=.002; vals=[]
72 for _ in range(500):
73 vals.append(free_energy(x,xs)); x=rk4_birth_death(x,u0,k,dt)
74 decay=np.polyfit(np.arange(100,450)*dt,np.log(np.maximum(vals[100:450],1e-30)),1)[0]
75 ratios=Vs/amps**2
76 # Direct mass-action formula check: X -> 2X and X -> empty gives u*x - k*x.
77 block=PositiveMassAction([[0],[1]], [[1],[0]], dt=.1)
78 field_error=abs(float(block.field(np.array([1.3]),np.array([u0,k]))[0])-(u0-k*1.3))
79 return {'euler_boundary_observed':float(lo),'euler_boundary_predicted':2.0,
80 'steady_energy_log_slope_observed':float(slope),'steady_energy_log_slope_predicted':2.0,
81 'small_amplitude_V_over_a2_observed':float(np.mean(ratios[:3])),
82 'small_amplitude_V_over_a2_predicted':.5,
83 'nominal_log_V_rate_observed':float(decay),'nominal_log_V_rate_predicted':-2.0,
84 'mass_action_field_absolute_error':float(field_error),
85 'amplitudes':amps.tolist(),'steady_V':Vs.tolist(),'V_over_a2':ratios.tolist()}
86
87
88def sequence_demo(seed=7):
89 rng=np.random.default_rng(seed); T=50; nseq=60
90 X=rng.uniform(-1,1,(nseq,T)); Y=np.zeros_like(X)
91 for b in range(nseq):
92 y=0.
93 for t in range(T): y=.9*y+.1*X[b,t]; Y[b,t]=y
94 sp=lambda z: np.logaddexp(0,z)
95 def pos_loss(p):
96 a,b,w,c=p; x=np.ones(nseq); se=0.
97 for t in range(T):
98 x += .2*(sp(a*X[:,t]+b)-x)
99 se += np.sum((w*x+c-Y[:,t])**2)
100 return se/(nseq*T)
101 def tanh_loss(p):
102 a,b,w,c=p; h=np.zeros(nseq); se=0.
103 for t in range(T):
104 h=np.tanh(a*X[:,t]+b*h)
105 se += np.sum((w*h+c-Y[:,t])**2)
106 return se/(nseq*T)
107 best=(1e9,None); best2=(1e9,None)
108 for _ in range(500):
109 p=rng.normal(0,1,4); p[2]=rng.normal(0,2); p[3]=rng.normal(0,.3)
110 z=pos_loss(p)
111 if z<best[0]: best=(z,p.copy())
112 p=rng.normal(0,1,4); z=tanh_loss(p)
113 if z<best2[0]: best2=(z,p.copy())
114 def long_eval(p,H,positive):
115 a,b,w,c=p; errs=[]
116 for row in X:
117 state=1. if positive else 0.; yy=0.
118 for t in range(H):
119 q=row[t%T]+.15*rng.normal(); yy=.9*yy+.1*q
120 state=(state+.2*(sp(a*q+b)-state)) if positive else np.tanh(a*q+b*state)
121 errs.append((w*state+c-yy)**2)
122 return float(np.mean(errs))
123 return {'positive_train_mse':float(best[0]),'tanh_train_mse':float(best2[0]),
124 'positive_long_noisy_mse':long_eval(best[1],400,True),
125 'tanh_long_noisy_mse':long_eval(best2[1],400,False),
126 'positive_params':best[1].tolist(),'tanh_params':best2[1].tolist()}
127
128
129if __name__ == '__main__':
130 out={'verification':verify(),'sequence_demo':sequence_demo()}
131 Path('results.json').write_text(json.dumps(out,indent=2))
132 print(json.dumps(out,indent=2))