Spectral-Edge Criticality Controller / bench_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import sys, json, math, time
 2import numpy as np
 3import torch
 4from torch import nn
 5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 6import bench
 7
 8class ControlledGRU(nn.Module):
 9    def __init__(self, out_dim, controlled=False, target=0.9, alpha=0.25):
10        super().__init__()
11        self.hdim, self.indim = 64, 3
12        self.wih = nn.Parameter(torch.empty(192, 3)); self.whh = nn.Parameter(torch.empty(192, 64))
13        self.bih = nn.Parameter(torch.zeros(192)); self.bhh = nn.Parameter(torch.zeros(192))
14        self.head = nn.Linear(64, out_dim)
15        nn.init.xavier_uniform_(self.wih); nn.init.orthogonal_(self.whh)
16        self.controlled, self.target, self.alpha = controlled, target, alpha
17        self.gain = 1.0; self.edges = []; self.gains = []
18
19    @torch.no_grad()
20    def estimate_edge(self, h, x):
21        # Local Jacobian of the realized GRU map, approximated by diagonal gate
22        # sensitivity times recurrent matrix; power iteration estimates its 2-norm.
23        q = torch.randn(h.shape[0], self.hdim, device=h.device)
24        q = q / q.norm(dim=1, keepdim=True).clamp_min(1e-8)
25        z = x @ self.wih.T + self.bih + h @ self.whh.T + self.bhh
26        _, u, n = z.chunk(3, dim=1)
27        u = torch.sigmoid(u); n = torch.tanh(n)
28        d = (1-u) + u * (1-n*n)
29        W = self.whh[128:192]
30        for _ in range(4):
31            q = d * (q @ W.T)
32            q = q / q.norm(dim=1, keepdim=True).clamp_min(1e-8)
33        return float((d * (q @ W.T)).norm(dim=1).mean().cpu())
34
35    def forward(self, x):
36        seq = x.view(x.shape[0], -1, 3); h = torch.zeros(x.shape[0], self.hdim, device=x.device)
37        self.edges=[]; self.gains=[]
38        for xt in seq.transpose(0,1):
39            z = xt @ self.wih.T + self.bih + (self.gain*h) @ self.whh.T + self.bhh
40            r,u,n = z.chunk(3, dim=1); r=torch.sigmoid(r); u=torch.sigmoid(u); n=torch.tanh(n + r*(h @ self.whh[128:192].T))
41            h = (1-u)*h + u*n
42            if self.controlled:
43                edge=self.estimate_edge(h.detach(), xt.detach())
44                self.gain *= math.exp(self.alpha*(self.target-self.gain*edge))
45                self.gain=float(np.clip(self.gain, 0.05, 2.0))
46            else: edge=self.estimate_edge(h.detach(), xt.detach())
47            self.edges.append(edge*self.gain); self.gains.append(self.gain)
48        return self.head(h)
49
50def train_one(track, seed, controlled, cfg):
51    np.random.seed(seed); torch.manual_seed(seed)
52    ds=bench.get_dataset(track, seed=seed, n_train=400, n_test=200)
53    model=ControlledGRU(ds['out_dim'], controlled, cfg.get('target',.9), cfg.get('alpha',.25))
54    model, metric, hist=bench.train_model(model, ds, epochs=cfg['epochs'], lr=cfg['lr'], batch=128)
55    sig={'edge_mean':float(np.mean(model.edges)) if model and model.edges else float('nan'), 'gain_final':float(model.gain) if model else float('nan')}
56    return float(metric), sig
57
58def eval_cfg(controlled, cfg, seeds):
59    vals=[]; sig=[]
60    for s in seeds:
61        v,z=train_one('dynamics',s,controlled,cfg); vals.append(v); sig.append(z)
62    return {'per_seed':vals,'mean':float(np.mean(vals)),'signatures':sig}
63
64def main():
65    # Same union of learning rates and equal 3-point search on both sides.
66    seeds=list(range(8)); configs=[{'lr':1e-3,'epochs':12},{'lr':3e-3,'epochs':12},{'lr':6e-3,'epochs':12}]
67    base_sweep=[]; idea_sweep=[]
68    for c in configs:
69        b=eval_cfg(False,c,seeds[:4]); base_sweep.append({'cfg':c,'mean':b['mean']})
70        i=eval_cfg(True,c,seeds[:4]); idea_sweep.append({'cfg':c,'mean':i['mean']})
71    bc=min(configs,key=lambda c: next(x['mean'] for x in base_sweep if x['cfg']==c))
72    ic=min(configs,key=lambda c: next(x['mean'] for x in idea_sweep if x['cfg']==c))
73    baseline=eval_cfg(False,bc,seeds); idea=eval_cfg(True,ic,seeds)
74    diffs=[a-b for a,b in zip(idea['per_seed'],baseline['per_seed'])]
75    p=bench.permutation_pvalue(diffs)
76    pred=ic.get('target',.9); observed=float(np.mean([x['edge_mean'] for x in idea['signatures']]))
77    report={'bench_version':1,'track':'dynamics','model':'rnn_small','metric_direction':'lower is better','n_seeds':8,
78      'baseline':{'best_cfg':bc,'sweep':base_sweep,'full':baseline},'idea':{'best_cfg':ic,'sweep':idea_sweep,'per_seed':idea['per_seed'],'mean':idea['mean'],'signatures':idea['signatures']},
79      'comparison':{'delta_mean':float(np.mean(diffs)),'per_seed_diffs':diffs,'idea_wins':sum(d<0 for d in diffs),'n_pairs':8,'p_value':p,'verdict':'idea better (significant)' if np.mean(diffs)<0 and p<.05 else 'no significant win'},
80      'mechanism_signature':{'prediction':'controlled effective edge near target below 1','predicted_target':pred,'observed_edge_mean':observed,'relative_error':abs(observed-pred)/pred,'confirmed':bool(abs(observed-pred)/pred<.2)}}
81    with open('bench_report.json','w') as f: json.dump(report,f,indent=2)
82    print(json.dumps(report,indent=2))
83if __name__=='__main__': main()