import sys, json, random from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, train_model, sweep_baseline, make_report, compare_results SEEDS = tuple(range(8)) SWEEP = (0,1,2,3) # Equal union of learning rates; retention knobs are swept on both sides. LRS = (1e-3, 3e-3, 6e-3) RETENTIONS = (0.5, 0.8, 1.0) GATE_BIASES = (-1.0, 0.0, 1.0) EPOCHS = 10 NTR, NTE = 800, 300 class PersistentWorkspace(nn.Module): def __init__(self, gated=False, retention=1.0, gate_bias=0.0, hidden=32, write=8): super().__init__() self.hidden, self.write, self.gated = hidden, write, gated self.enc = nn.Linear(1, write) self.trans = nn.Sequential(nn.Linear(hidden, hidden), nn.Tanh()) self.read = nn.Linear(hidden, 1) self.retention = retention if gated: # Gate depends on current state and encoded observation, as proposed. self.gate = nn.Linear(hidden + write, hidden) nn.init.constant_(self.gate.bias, gate_bias) else: self.gate = None def forward(self, x, return_state=False): b, t = x.shape s = torch.zeros(b, self.hidden, device=x.device, dtype=x.dtype) for k in range(t): z = self.enc(x[:, k:k+1]) # Explicit write port; workspace coordinates persist. s = torch.cat((z, s[:, self.write:]), dim=1) s = self.trans(s) if k != t - 1: if self.gated: g = torch.sigmoid(self.gate(torch.cat((s, z), dim=1))) s = g * s else: s = self.retention * s y = self.read(s) if return_state: return y, s return y def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) def one_train(cfg, seed, gated, keep_model=False): seed_all(seed) ds = get_dataset('sequence', seed, n_train=NTR, n_test=NTE) model = PersistentWorkspace(gated=gated, retention=cfg.get('retention',1.0), gate_bias=cfg.get('gate_bias',0.0)) net, metric, hist = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'], batch=128, weight_decay=0.0, log=lambda *_: None) if keep_model: return float(metric), net, ds return float(metric) def evaluator(cfg, gated): return lambda seed: one_train(cfg, seed, gated) def main(): baseline_grid = [{'lr':lr, 'retention':r} for lr in LRS for r in RETENTIONS] idea_grid = [{'lr':lr, 'gate_bias':b} for lr in LRS for b in GATE_BIASES] base = sweep_baseline(lambda c: evaluator(c, False), baseline_grid, seeds=SWEEP) # Idea grid is evaluated on all paired seeds; best is selected using the same # four-seed development protocol, then independently compared on eight seeds. idea_trials = [] for cfg in idea_grid: vals = [one_train(cfg, s, True) for s in SWEEP] idea_trials.append({'cfg':cfg, 'mean':float(np.mean(vals))}) best_idea_cfg = min(idea_trials, key=lambda z:z['mean'])['cfg'] idea_full = {'best_cfg':best_idea_cfg, 'sweep':idea_trials, 'per_seed':[one_train(best_idea_cfg,s,True) for s in SEEDS]} idea_full.update(mean=float(np.mean(idea_full['per_seed'])), std=float(np.std(idea_full['per_seed'])), n=8) extra = mechanism_signature(base['best_cfg'], best_idea_cfg) report = make_report('sequence', 'persistent_workspace_shared_transition', base, idea_full, extra) Path('bench_report.json').write_text(json.dumps(report, indent=2)) print(json.dumps(report, indent=2)) def mechanism_signature(base_cfg, idea_cfg): # Measure trained models, not an analytic toy: inject a state and observe # its norm after n no-input recurrence/retention steps. Compare to the # gate's observed mean multiplier prediction g^n. seed = 0 _, base_net, ds = one_train(base_cfg, seed, False, True) _, idea_net, _ = one_train(idea_cfg, seed, True, True) device = next(idea_net.parameters()).device with torch.no_grad(): x = ds['xte'][:64].to(device) _, sb = base_net(x, True); _, si = idea_net(x, True) # Use a repeated zero observation, retaining the trained state. z = torch.zeros_like(x[:, :1]) def evolve(net, s, gated): norms=[s.norm(dim=1).mean().item()] for _ in range(5): zz=net.enc(z); s=torch.cat((zz,s[:,net.write:]),1); s=net.trans(s) if gated: s=torch.sigmoid(net.gate(torch.cat((s,zz),1)))*s else: s=net.retention*s norms.append(s.norm(dim=1).mean().item()) return norms bn=evolve(base_net,sb,False); inn=evolve(idea_net,si,True) observed=float(np.mean([inn[i+1]/(inn[i]+1e-8) for i in range(5)])) predicted=float(np.mean([torch.sigmoid(idea_net.gate(torch.cat((idea_net.trans(torch.cat((idea_net.enc(z),si[:,idea_net.write:]),1)), idea_net.enc(z)),1))).mean().item() for _ in [0]])) return {'quantity':'trained-state retention multiplier per macrostep', 'predicted_mean_gate':predicted, 'observed_mean_norm_ratio':observed, 'baseline_norm_ratios': [bn[i+1]/(bn[i]+1e-8) for i in range(5)], 'idea_norms':inn, 'confirmed': bool(abs(predicted-observed) < 0.15)} if __name__ == '__main__': main()