Persistent Workspace for Online Adaptation / stage2_bench.py
Mechanism confirmed, baseline not beaten
1import sys, json, random
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, train_model, sweep_baseline, make_report, compare_results
8
9SEEDS = tuple(range(8))
10SWEEP = (0,1,2,3)
11# Equal union of learning rates; retention knobs are swept on both sides.
12LRS = (1e-3, 3e-3, 6e-3)
13RETENTIONS = (0.5, 0.8, 1.0)
14GATE_BIASES = (-1.0, 0.0, 1.0)
15EPOCHS = 10
16NTR, NTE = 800, 300
17
18class PersistentWorkspace(nn.Module):
19 def __init__(self, gated=False, retention=1.0, gate_bias=0.0, hidden=32, write=8):
20 super().__init__()
21 self.hidden, self.write, self.gated = hidden, write, gated
22 self.enc = nn.Linear(1, write)
23 self.trans = nn.Sequential(nn.Linear(hidden, hidden), nn.Tanh())
24 self.read = nn.Linear(hidden, 1)
25 self.retention = retention
26 if gated:
27 # Gate depends on current state and encoded observation, as proposed.
28 self.gate = nn.Linear(hidden + write, hidden)
29 nn.init.constant_(self.gate.bias, gate_bias)
30 else:
31 self.gate = None
32
33 def forward(self, x, return_state=False):
34 b, t = x.shape
35 s = torch.zeros(b, self.hidden, device=x.device, dtype=x.dtype)
36 for k in range(t):
37 z = self.enc(x[:, k:k+1])
38 # Explicit write port; workspace coordinates persist.
39 s = torch.cat((z, s[:, self.write:]), dim=1)
40 s = self.trans(s)
41 if k != t - 1:
42 if self.gated:
43 g = torch.sigmoid(self.gate(torch.cat((s, z), dim=1)))
44 s = g * s
45 else:
46 s = self.retention * s
47 y = self.read(s)
48 if return_state:
49 return y, s
50 return y
51
52def seed_all(seed):
53 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
54
55def one_train(cfg, seed, gated, keep_model=False):
56 seed_all(seed)
57 ds = get_dataset('sequence', seed, n_train=NTR, n_test=NTE)
58 model = PersistentWorkspace(gated=gated, retention=cfg.get('retention',1.0),
59 gate_bias=cfg.get('gate_bias',0.0))
60 net, metric, hist = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'],
61 batch=128, weight_decay=0.0, log=lambda *_: None)
62 if keep_model:
63 return float(metric), net, ds
64 return float(metric)
65
66def evaluator(cfg, gated):
67 return lambda seed: one_train(cfg, seed, gated)
68
69def main():
70 baseline_grid = [{'lr':lr, 'retention':r} for lr in LRS for r in RETENTIONS]
71 idea_grid = [{'lr':lr, 'gate_bias':b} for lr in LRS for b in GATE_BIASES]
72 base = sweep_baseline(lambda c: evaluator(c, False), baseline_grid, seeds=SWEEP)
73 # Idea grid is evaluated on all paired seeds; best is selected using the same
74 # four-seed development protocol, then independently compared on eight seeds.
75 idea_trials = []
76 for cfg in idea_grid:
77 vals = [one_train(cfg, s, True) for s in SWEEP]
78 idea_trials.append({'cfg':cfg, 'mean':float(np.mean(vals))})
79 best_idea_cfg = min(idea_trials, key=lambda z:z['mean'])['cfg']
80 idea_full = {'best_cfg':best_idea_cfg, 'sweep':idea_trials,
81 'per_seed':[one_train(best_idea_cfg,s,True) for s in SEEDS]}
82 idea_full.update(mean=float(np.mean(idea_full['per_seed'])), std=float(np.std(idea_full['per_seed'])), n=8)
83 extra = mechanism_signature(base['best_cfg'], best_idea_cfg)
84 report = make_report('sequence', 'persistent_workspace_shared_transition', base, idea_full, extra)
85 Path('bench_report.json').write_text(json.dumps(report, indent=2))
86 print(json.dumps(report, indent=2))
87
88def mechanism_signature(base_cfg, idea_cfg):
89 # Measure trained models, not an analytic toy: inject a state and observe
90 # its norm after n no-input recurrence/retention steps. Compare to the
91 # gate's observed mean multiplier prediction g^n.
92 seed = 0
93 _, base_net, ds = one_train(base_cfg, seed, False, True)
94 _, idea_net, _ = one_train(idea_cfg, seed, True, True)
95 device = next(idea_net.parameters()).device
96 with torch.no_grad():
97 x = ds['xte'][:64].to(device)
98 _, sb = base_net(x, True); _, si = idea_net(x, True)
99 # Use a repeated zero observation, retaining the trained state.
100 z = torch.zeros_like(x[:, :1])
101 def evolve(net, s, gated):
102 norms=[s.norm(dim=1).mean().item()]
103 for _ in range(5):
104 zz=net.enc(z); s=torch.cat((zz,s[:,net.write:]),1); s=net.trans(s)
105 if gated: s=torch.sigmoid(net.gate(torch.cat((s,zz),1)))*s
106 else: s=net.retention*s
107 norms.append(s.norm(dim=1).mean().item())
108 return norms
109 bn=evolve(base_net,sb,False); inn=evolve(idea_net,si,True)
110 observed=float(np.mean([inn[i+1]/(inn[i]+1e-8) for i in range(5)]))
111 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]]))
112 return {'quantity':'trained-state retention multiplier per macrostep',
113 'predicted_mean_gate':predicted, 'observed_mean_norm_ratio':observed,
114 'baseline_norm_ratios': [bn[i+1]/(bn[i]+1e-8) for i in range(5)],
115 'idea_norms':inn, 'confirmed': bool(abs(predicted-observed) < 0.15)}
116
117if __name__ == '__main__': main()