import json, math, random from pathlib import Path import numpy as np import torch from torch import nn SEED = 1355 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(4) class SpectralMemoryLift(nn.Module): """Small batched implementation of q'=Aq+Bz+Fu, z'=Cq+Dz.""" def __init__(self, d=1, h=1, experts=4, taus=(2., 8., 32., 128.)): super().__init__(); self.d=d; self.h=h; self.E=experts self.A=nn.Parameter(.05*torch.randn(experts,d,d)) self.B=nn.Parameter(.1*torch.randn(experts,d,h)) self.C=nn.Parameter(.1*torch.randn(experts,h,d)) self.F=nn.Parameter(.1*torch.randn(experts,d,d)) self.Dhat=nn.Parameter(torch.randn(experts,h,h)) # Distinct initial memory rates; bounded through sigmoid. rates = torch.tensor([math.exp(-1./t) for t in taus[:experts]]) logits = torch.logit(rates.clamp(.01,.99)) self.rate_logits=nn.Parameter(logits[:,None].repeat(1,h)) self.read=nn.Parameter(.1*torch.randn(experts,d+h)) self.gate=nn.Linear(d+experts*(d+h), experts) self.taus=torch.tensor(list(taus[:experts]), dtype=torch.float32) def D(self): raw=self.Dhat/(self.Dhat.norm(dim=(-2,-1),keepdim=True)+1e-8) return raw*torch.sigmoid(self.rate_logits)[:,:,None] def forward(self, u, return_states=False): # u: [batch,time,d]; q and z are reset per sequence. b,T,_=u.shape; q=torch.zeros(b,self.E,self.d,device=u.device) z=torch.zeros(b,self.E,self.h,device=u.device); D=self.D() ys=[]; states=[] for t in range(T): x=u[:,t] qn=torch.einsum('edk,bek->bed',self.A,q)+torch.einsum('edh,beh->bed',self.B,z)+torch.einsum('edk,bk->bed',self.F,x) zn=torch.einsum('ehd,bed->beh',self.C,q)+torch.einsum('ehi,bei->beh',D,z) vals=torch.cat([qn,zn],-1) logits=self.gate(torch.cat([x,vals.reshape(b,-1)],-1)) # stop-gradient slow-mode feature, as proposed logits=logits + .05*(-1./torch.log(self.taus.clamp_min(1.01))).to(x.device) p=torch.softmax(logits,-1) pred=(p.unsqueeze(-1)*torch.einsum('eK,beK->be',self.read,vals).unsqueeze(-1)).sum(1) ys.append(pred); q,z=qn,zn; states.append((q,z)) out=torch.stack(ys,1) return (out,states) if return_states else out class SingleExponential(nn.Module): """Baseline: one stable scalar memory and linear readout.""" def __init__(self): super().__init__(); self.logit=nn.Parameter(torch.tensor(0.0)); self.a=nn.Parameter(torch.tensor(.1)); self.b=nn.Parameter(torch.tensor(.1)); self.c=nn.Parameter(torch.tensor(.1)) def forward(self,u): b,T,_=u.shape; z=torch.zeros(b,1,device=u.device); ys=[] for t in range(T): z=torch.sigmoid(self.logit)*z+self.a*u[:,t] ys.append(self.b*z+self.c*u[:,t]) return torch.stack(ys,1) def exact_kernel_check(): # For z_{n+1}=Cq_n+Dz_n, contribution of q_n to q_{n+k} through Bz is BD^(k-1)C. A=np.array([[.2]]); B=np.array([[.7]]); C=np.array([[.4]]); D=np.array([[.8]]) q=1.; z=0.; observed=[] # impulse q_0, then suppress direct q recurrence and inject only its memory path for k in range(1,8): z=C*q + D*z if k==1 else D*z observed.append(float(B@z)); q=0. predicted=[float(B@np.linalg.matrix_power(D,k-1)@C) for k in range(1,8)] return max(abs(np.array(observed)-predicted)), observed, predicted def stability_sweep(): # Scale a fixed 2x2 transition M by gamma. Boundary is gamma*rho(M0)=1. M0=np.array([[.35,.25],[-.15,.70]]) rho=max(abs(np.linalg.eigvals(M0))); boundary=1/rho rows=[] for gamma in np.linspace(.4*boundary,1.6*boundary,13): M=gamma*M0; eig=np.linalg.eigvals(M); r=max(abs(eig)) x=np.array([1.,-.3]); norms=[] for _ in range(80): x=M@x; norms.append(np.linalg.norm(x)) # asymptotic ratio, robustly measured over last 20 steps ratio=float((norms[-1]/norms[-21])**(1/20)) rows.append({'gamma':float(gamma),'rho':float(r),'ratio':ratio,'stable':bool(norms[-1]=2: target[:,t]+=0.55*x[:,t-2] if t>=24: target[:,t]+=0.30*x[:,t-24] if t>=64: target[:,t]+=0.15*x[:,t-64] split=192; models=[SingleExponential().to(device),SpectralMemoryLift(experts=4,taus=(2,8,32,128)).to(device)] results=[] for model in models: opt=torch.optim.Adam(model.parameters(),lr=0.02) for step in range(500): pred=model(x[:split]); loss=((pred-target[:split])**2).mean() opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(),5.); opt.step() with torch.no_grad(): tr=float(((model(x[:split])-target[:split])**2).mean()); va=float(((model(x[split:])-target[split:])**2).mean()) out=model(x[split:]); state_norm=None if isinstance(model,SpectralMemoryLift): _,states=model(x[split:],True); state_norm=float(max(max(q.norm(dim=-1).max().item(),z.norm(dim=-1).max().item()) for q,z in states)) else: state_norm=float(out.abs().max()) results.append({'model':type(model).__name__,'train_mse':tr,'validation_mse':va,'max_state_or_output_norm':state_norm}) if device=='cuda': torch.cuda.empty_cache() return results except Exception as e: # Required safe fallback if CUDA/runtime allocation fails. torch.cuda.empty_cache() if torch.cuda.is_available() else None return [{'error':str(e),'fallback':'CPU unavailable in this run'}] def main(): kernel_err,obs,pred=exact_kernel_check() result={'seed':SEED,'kernel_check':{'max_abs_error':kernel_err,'observed':obs,'predicted':pred},'stability_check':stability_sweep(),'timescale_check':timescale_sweep(),'mini_experiment':mini_experiment()} Path('results.json').write_text(json.dumps(result,indent=2)) print(json.dumps(result,indent=2)) if __name__=='__main__': main()