Fading-Memory Habituation Gate / stage2_bench.py
Unverified
1import sys, json, random, math
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, evaluate, sweep_baseline, make_report
8
9SEEDS = tuple(range(8))
10LR_GRID = [0.001, 0.003, 0.006]
11EPOCHS = 8
12IDEA_GRID = [(0.85, 0.5), (0.90, 0.5), (0.90, 1.0)]
13
14
15def seed_all(seed):
16 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
17 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
18
19
20class HabituationTransformer(nn.Module):
21 def __init__(self, win, out_dim, rho=0.9, beta=0.5, enabled=True):
22 super().__init__()
23 d = 64
24 self.win, self.rho, self.beta, self.enabled = win, float(rho), float(beta), enabled
25 self.inp = nn.Linear(1, d)
26 self.pos = nn.Parameter(torch.zeros(1, win, d))
27 nn.init.normal_(self.pos, std=.02)
28 layer = nn.TransformerEncoderLayer(d, nhead=2, dim_feedforward=128,
29 batch_first=True, dropout=0.0)
30 self.enc = nn.TransformerEncoder(layer, 2)
31 self.head = nn.Linear(win*d, out_dim)
32 self.last_states = None
33 self.last_gains = None
34
35 def forward(self, x, capture=False):
36 v = self.inp(x.unsqueeze(-1)) + self.pos[:, :x.shape[1]]
37 if not self.enabled:
38 h = self.enc(v)
39 return self.head(h.reshape(x.shape[0], -1))
40 # Scalar per-token stimulation, with stop-gradient to isolate habituation dynamics.
41 a = torch.zeros(v.shape[0], device=v.device, dtype=v.dtype)
42 hs, states, gains = [], [], []
43 for t in range(v.shape[1]):
44 vt = v[:, t]
45 s = torch.sqrt((vt.detach() ** 2).mean(dim=-1) + 1e-8)
46 a = self.rho * a + (1.0 - self.rho) * s
47 g = 1.0 / (1.0 + self.beta * a)
48 hs.append(vt * g[:, None])
49 if capture:
50 states.append(a.detach()); gains.append(g.detach())
51 h = self.enc(torch.stack(hs, dim=1))
52 if capture:
53 self.last_states = torch.stack(states, dim=1)
54 self.last_gains = torch.stack(gains, dim=1)
55 return self.head(h.reshape(x.shape[0], -1))
56
57
58def train_one(seed, cfg, idea, capture=False):
59 seed_all(seed)
60 ds = get_dataset('sequence', seed, n_train=400, n_test=200)
61 win = ds['input_shape'][0]
62 net = HabituationTransformer(win, ds['out_dim'], cfg.get('rho', .9),
63 cfg.get('beta', .5), enabled=idea)
64 net, metric, history = train_model(net, ds, epochs=cfg['epochs'], lr=cfg['lr'], batch=128,
65 log=lambda *_: None)
66 if not capture:
67 return float(metric)
68 net.eval()
69 dev = next(net.parameters()).device
70 xte = ds['xte'].to(dev)
71 with torch.no_grad():
72 pred = net(xte, capture=True)
73 gains = net.last_gains.cpu().numpy() if net.last_gains is not None else np.ones((len(ds['xte']), win))
74 states = net.last_states.cpu().numpy() if net.last_states is not None else np.zeros_like(gains)
75 # Signature is measured on the trained network's actual test-window representations.
76 observed_final = float(gains[:, -1].mean())
77 observed_initial = float(gains[:, 0].mean())
78 stim = float(states[:, -1].mean())
79 predicted_final = 1.0 / (1.0 + cfg.get('beta', 0.0) * stim) if idea else 1.0
80 # Fit temporal state pole using the first test sample's measured state trajectory.
81 if idea and states.shape[1] > 2:
82 y = states[0, 1:]; x = states[0, :-1]
83 slope = float(np.dot(x, y) / max(np.dot(x, x), 1e-12))
84 else:
85 slope = 0.0
86 return {'metric': float(metric), 'initial_gain': observed_initial,
87 'final_gain': observed_final, 'mean_state': stim,
88 'predicted_steady_gain': predicted_final, 'fitted_state_pole': slope,
89 'rho': cfg.get('rho', 0.0), 'beta': cfg.get('beta', 0.0)}
90
91
92def math_check():
93 rows=[]
94 for rho,beta in [(0.85,.5),(.9,.5),(.9,1.0)]:
95 a=0.0
96 for _ in range(300): a=rho*a+(1-rho)*1.7
97 gain=1/(1+beta*a); pred=1/(1+beta*1.7)
98 rec=[]; z=a
99 for _ in range(100): z=rho*z; rec.append(z)
100 half=next((i+1 for i,q in enumerate(rec) if q<=a/2),None)
101 ph=math.ceil(math.log(.5)/math.log(rho))
102 rows.append({'rho':rho,'steady_gain_abs_error':abs(gain-pred),
103 'half_observed':half,'half_predicted':ph})
104 return {'rows':rows,'confirmed':all(r['steady_gain_abs_error']<1e-10 and r['half_observed']==r['half_predicted'] for r in rows)}
105
106
107def main():
108 # Baseline sweep includes every LR used by the idea side, satisfying union parity.
109 baseline_grid=[{'lr':lr,'epochs':EPOCHS,'rho':.0,'beta':0.0} for lr in LR_GRID]
110 base=sweep_baseline(lambda cfg: (lambda seed: train_one(seed,cfg,False)), baseline_grid, seeds=SEEDS)
111 best=base['best_cfg']
112 idea_cfgs=[]
113 for lr in LR_GRID:
114 rho,beta=IDEA_GRID[len(idea_cfgs)%len(IDEA_GRID)]
115 idea_cfgs.append({'lr':lr,'epochs':EPOCHS,'rho':rho,'beta':beta})
116 idea_runs=[]
117 for cfg in idea_cfgs:
118 result=evaluate(lambda seed,cfg=cfg: train_one(seed,cfg,True), seeds=SEEDS)
119 idea_runs.append((cfg,result))
120 idea_cfg, idea = min(idea_runs, key=lambda z:z[1]['mean'])
121 sig=train_one(0, idea_cfg, True, capture=True)
122 confirmed=(abs(sig['final_gain']-sig['predicted_steady_gain']) < .08 and
123 abs(sig['fitted_state_pole']-idea_cfg['rho']) < .12)
124 report=make_report('sequence','transformer_tiny',base,idea,extra={
125 'track_rationale':'Sequence forecast contains multi-token temporal correlations and matches the proposed per-token fading-memory transformer gate.',
126 'observed_best_cfg':idea_cfg,
127 'idea_sweep':[{'cfg':c,'result':r} for c,r in idea_runs],
128 'mechanism_signature':{'trained_model_seed':0, **sig, 'confirmed':bool(confirmed)},
129 'math_check':math_check(),
130 'protocol':{'seeds':list(SEEDS),'epochs':EPOCHS,'baseline_lr_grid':LR_GRID,'idea_lr_grid':LR_GRID}
131 })
132 Path('bench_report.json').write_text(json.dumps(report,indent=2))
133 print(json.dumps(report,indent=2))
134
135if __name__=='__main__': main()