ISS-Gated Positive Neural State Module / iss_experiment.py

Failed on benchmark

Raw ⬇ ZIP
  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))