Spectral Memory-Lift Ensemble / spectral_memory_lift.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1import json, math, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5from torch import nn
  6
  7SEED = 1355
  8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
  9torch.set_num_threads(4)
 10
 11class SpectralMemoryLift(nn.Module):
 12    """Small batched implementation of q'=Aq+Bz+Fu, z'=Cq+Dz."""
 13    def __init__(self, d=1, h=1, experts=4, taus=(2., 8., 32., 128.)):
 14        super().__init__(); self.d=d; self.h=h; self.E=experts
 15        self.A=nn.Parameter(.05*torch.randn(experts,d,d))
 16        self.B=nn.Parameter(.1*torch.randn(experts,d,h))
 17        self.C=nn.Parameter(.1*torch.randn(experts,h,d))
 18        self.F=nn.Parameter(.1*torch.randn(experts,d,d))
 19        self.Dhat=nn.Parameter(torch.randn(experts,h,h))
 20        # Distinct initial memory rates; bounded through sigmoid.
 21        rates = torch.tensor([math.exp(-1./t) for t in taus[:experts]])
 22        logits = torch.logit(rates.clamp(.01,.99))
 23        self.rate_logits=nn.Parameter(logits[:,None].repeat(1,h))
 24        self.read=nn.Parameter(.1*torch.randn(experts,d+h))
 25        self.gate=nn.Linear(d+experts*(d+h), experts)
 26        self.taus=torch.tensor(list(taus[:experts]), dtype=torch.float32)
 27    def D(self):
 28        raw=self.Dhat/(self.Dhat.norm(dim=(-2,-1),keepdim=True)+1e-8)
 29        return raw*torch.sigmoid(self.rate_logits)[:,:,None]
 30    def forward(self, u, return_states=False):
 31        # u: [batch,time,d]; q and z are reset per sequence.
 32        b,T,_=u.shape; q=torch.zeros(b,self.E,self.d,device=u.device)
 33        z=torch.zeros(b,self.E,self.h,device=u.device); D=self.D()
 34        ys=[]; states=[]
 35        for t in range(T):
 36            x=u[:,t]
 37            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)
 38            zn=torch.einsum('ehd,bed->beh',self.C,q)+torch.einsum('ehi,bei->beh',D,z)
 39            vals=torch.cat([qn,zn],-1)
 40            logits=self.gate(torch.cat([x,vals.reshape(b,-1)],-1))
 41            # stop-gradient slow-mode feature, as proposed
 42            logits=logits + .05*(-1./torch.log(self.taus.clamp_min(1.01))).to(x.device)
 43            p=torch.softmax(logits,-1)
 44            pred=(p.unsqueeze(-1)*torch.einsum('eK,beK->be',self.read,vals).unsqueeze(-1)).sum(1)
 45            ys.append(pred); q,z=qn,zn; states.append((q,z))
 46        out=torch.stack(ys,1)
 47        return (out,states) if return_states else out
 48
 49class SingleExponential(nn.Module):
 50    """Baseline: one stable scalar memory and linear readout."""
 51    def __init__(self):
 52        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))
 53    def forward(self,u):
 54        b,T,_=u.shape; z=torch.zeros(b,1,device=u.device); ys=[]
 55        for t in range(T):
 56            z=torch.sigmoid(self.logit)*z+self.a*u[:,t]
 57            ys.append(self.b*z+self.c*u[:,t])
 58        return torch.stack(ys,1)
 59
 60def exact_kernel_check():
 61    # For z_{n+1}=Cq_n+Dz_n, contribution of q_n to q_{n+k} through Bz is BD^(k-1)C.
 62    A=np.array([[.2]]); B=np.array([[.7]]); C=np.array([[.4]]); D=np.array([[.8]])
 63    q=1.; z=0.; observed=[]
 64    # impulse q_0, then suppress direct q recurrence and inject only its memory path
 65    for k in range(1,8):
 66        z=C*q + D*z if k==1 else D*z
 67        observed.append(float(B@z)); q=0.
 68    predicted=[float(B@np.linalg.matrix_power(D,k-1)@C) for k in range(1,8)]
 69    return max(abs(np.array(observed)-predicted)), observed, predicted
 70
 71def stability_sweep():
 72    # Scale a fixed 2x2 transition M by gamma. Boundary is gamma*rho(M0)=1.
 73    M0=np.array([[.35,.25],[-.15,.70]])
 74    rho=max(abs(np.linalg.eigvals(M0))); boundary=1/rho
 75    rows=[]
 76    for gamma in np.linspace(.4*boundary,1.6*boundary,13):
 77        M=gamma*M0; eig=np.linalg.eigvals(M); r=max(abs(eig))
 78        x=np.array([1.,-.3]); norms=[]
 79        for _ in range(80): x=M@x; norms.append(np.linalg.norm(x))
 80        # asymptotic ratio, robustly measured over last 20 steps
 81        ratio=float((norms[-1]/norms[-21])**(1/20))
 82        rows.append({'gamma':float(gamma),'rho':float(r),'ratio':ratio,'stable':bool(norms[-1]<norms[0])})
 83    transition=min(rows,key=lambda x:abs(x['rho']-1))
 84    return {'rho_M0':float(rho),'predicted_gamma_boundary':float(boundary),'rows':rows,'closest_observed':transition}
 85
 86def timescale_sweep():
 87    # A pure mode r has amplitude r^n and tau=-1/log(r); fit tau from log amplitude.
 88    rows=[]
 89    for r in [.5,.7,.8,.9,.97,.99]:
 90        tau=-1/math.log(r); n=np.arange(1,1000); amp=r**n
 91        fit=-1/np.polyfit(n,np.log(amp),1)[0]
 92        half=n[np.argmin(abs(amp-.5))]
 93        rows.append({'rho':r,'predicted_tau':tau,'fitted_tau':float(fit),'predicted_half_life':math.log(.5)/math.log(r),'observed_half_life':int(half)})
 94    return rows
 95
 96def mini_experiment():
 97    # Target has two separated exponential/delayed components, generated from input history.
 98    device='cuda' if torch.cuda.is_available() else 'cpu'
 99    try:
100        N,T=256,96; g=torch.Generator().manual_seed(SEED)
101        x=torch.randn(N,T,1,generator=g).to(device)
102        target=torch.zeros_like(x)
103        for t in range(T):
104            if t>=2: target[:,t]+=0.55*x[:,t-2]
105            if t>=24: target[:,t]+=0.30*x[:,t-24]
106            if t>=64: target[:,t]+=0.15*x[:,t-64]
107        split=192; models=[SingleExponential().to(device),SpectralMemoryLift(experts=4,taus=(2,8,32,128)).to(device)]
108        results=[]
109        for model in models:
110            opt=torch.optim.Adam(model.parameters(),lr=0.02)
111            for step in range(500):
112                pred=model(x[:split]); loss=((pred-target[:split])**2).mean()
113                opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(),5.); opt.step()
114            with torch.no_grad():
115                tr=float(((model(x[:split])-target[:split])**2).mean()); va=float(((model(x[split:])-target[split:])**2).mean())
116                out=model(x[split:]); state_norm=None
117                if isinstance(model,SpectralMemoryLift):
118                    _,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))
119                else: state_norm=float(out.abs().max())
120            results.append({'model':type(model).__name__,'train_mse':tr,'validation_mse':va,'max_state_or_output_norm':state_norm})
121        if device=='cuda': torch.cuda.empty_cache()
122        return results
123    except Exception as e:
124        # Required safe fallback if CUDA/runtime allocation fails.
125        torch.cuda.empty_cache() if torch.cuda.is_available() else None
126        return [{'error':str(e),'fallback':'CPU unavailable in this run'}]
127
128def main():
129    kernel_err,obs,pred=exact_kernel_check()
130    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()}
131    Path('results.json').write_text(json.dumps(result,indent=2))
132    print(json.dumps(result,indent=2))
133if __name__=='__main__': main()