import json, math, random from pathlib import Path import numpy as np SEED = 7 np.random.seed(SEED); random.seed(SEED) def retention_sweep(): # f(s)=s, reset=0, constant gate g => m_n=g^n*m_0 exactly. n = 20; rows = [] for g in [0.0, 0.25, 0.5, 0.8, 0.9, 0.99]: observed = float(g ** n) if 0 < g < 1: half_pred = math.log(0.5) / math.log(g) half_obs = next(k for k in range(101) if g ** k <= 0.5) elif g == 0: half_pred, half_obs = 1.0, 1 else: half_pred, half_obs = float('inf'), None rows.append({'g':g, 'steps':n, 'observed':observed, 'predicted':float(g**n), 'half_life_observed':half_obs, 'half_life_predicted':half_pred}) return rows def stability_sweep(): # f(s)=lambda*s then retention s+=g*s gives multiplier q=lambda*g. rows=[] for lam in [0.8, 1.0, 1.05, 1.2]: for g in [0.5, 0.8, 0.99, 1.0, 1.1]: q=lam*g; s=1.0 for _ in range(40): s*=q rows.append({'lambda':lam,'g':g,'multiplier':q,'final_abs':abs(s), 'observed_stable':abs(s)<=1.0, 'predicted_stable':q<=1.0}) return rows def overwrite_sweep(): # s+=g*s-+(1-g)*z, old=1,z=0 => overwrite amount 1-g. rows=[] for g in [0.0,.25,.5,.75,1.0]: out=g rows.append({'g':g,'observed_overwrite':1-out, 'predicted_overwrite':1-g,'resulting_state':out}) return rows def delayed_bit_train(): # Stream: bit, distractors, query. State is partitioned into a two-value # write port and persistent workspace. Only the query output is scored. try: import torch import torch.nn as nn torch.manual_seed(SEED) try: device='cuda' if torch.cuda.is_available() else 'cpu' # Small tensors only; fall back if CUDA is unavailable/unhealthy. torch.zeros(1, device=device) except Exception: device='cpu' class Workspace(nn.Module): def __init__(self, hidden=8, gated=True): super().__init__(); self.hidden=hidden; self.gated=gated self.enc=nn.Linear(2,hidden) self.trans=nn.Sequential(nn.Linear(hidden,hidden),nn.Tanh()) self.read=nn.Linear(hidden,1) self.gate=nn.Linear(2*hidden,hidden) if gated else None self.reset=nn.Parameter(torch.zeros(hidden)) def forward(self,x): b,time,_=x.shape # state starts with write-port zeros and workspace zeros s=torch.zeros(b,self.hidden,device=x.device) outs=[] for t in range(time): z=self.enc(x[:,t]) # Functional overwrite: write coordinates are replaced; # coordinates 2: are retained from the prior state. s=torch.cat([z[:,:2],s[:,2:]],dim=-1) s=self.trans(s) outs.append(self.read(s).squeeze(-1)) if self.gated: g=torch.sigmoid(self.gate(torch.cat([s,z],dim=-1))) s=g*s+(1-g)*self.reset return torch.stack(outs,dim=1) def make_batch(batch,gap): time=gap+2; x=torch.zeros(batch,time,2,device=device) bit=torch.randint(0,2,(batch,),device=device).float() x[:,0,0]=bit; x[:,-1,1]=1. return x,bit def run(gated,gap): torch.manual_seed(SEED+gap+int(gated)) model=Workspace(gated=gated).to(device) opt=torch.optim.Adam(model.parameters(),lr=.01) for _ in range(600): x,y=make_batch(64,gap); pred=model(x)[:,-1] loss=nn.functional.binary_cross_entropy_with_logits(pred,y) opt.zero_grad(); loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(),5); opt.step() with torch.no_grad(): x,y=make_batch(1024,gap) acc=((model(x)[:,-1]>0).float()==y).float().mean().item() return acc return {'device':device,'results':{ str(g):{'baseline_no_retention':run(False,g), 'gated_workspace':run(True,g)} for g in [2,8,16]}} except Exception as exc: return {'device':'cpu','error':repr(exc)} def main(): out={'retention_prediction':retention_sweep(), 'stability_prediction':stability_sweep(), 'overwrite_prediction':overwrite_sweep(), 'delayed_bit':delayed_bit_train()} Path('results.json').write_text(json.dumps(out,indent=2)); print(json.dumps(out,indent=2)) if __name__=='__main__': main()