Persistent Workspace for Online Adaptation / stage2_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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()