Fading-Memory Habituation Gate / stage2_bench.py

Unverified

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