import sys, json, math, time import numpy as np import torch from torch import nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') import bench class ControlledGRU(nn.Module): def __init__(self, out_dim, controlled=False, target=0.9, alpha=0.25): super().__init__() self.hdim, self.indim = 64, 3 self.wih = nn.Parameter(torch.empty(192, 3)); self.whh = nn.Parameter(torch.empty(192, 64)) self.bih = nn.Parameter(torch.zeros(192)); self.bhh = nn.Parameter(torch.zeros(192)) self.head = nn.Linear(64, out_dim) nn.init.xavier_uniform_(self.wih); nn.init.orthogonal_(self.whh) self.controlled, self.target, self.alpha = controlled, target, alpha self.gain = 1.0; self.edges = []; self.gains = [] @torch.no_grad() def estimate_edge(self, h, x): # Local Jacobian of the realized GRU map, approximated by diagonal gate # sensitivity times recurrent matrix; power iteration estimates its 2-norm. q = torch.randn(h.shape[0], self.hdim, device=h.device) q = q / q.norm(dim=1, keepdim=True).clamp_min(1e-8) z = x @ self.wih.T + self.bih + h @ self.whh.T + self.bhh _, u, n = z.chunk(3, dim=1) u = torch.sigmoid(u); n = torch.tanh(n) d = (1-u) + u * (1-n*n) W = self.whh[128:192] for _ in range(4): q = d * (q @ W.T) q = q / q.norm(dim=1, keepdim=True).clamp_min(1e-8) return float((d * (q @ W.T)).norm(dim=1).mean().cpu()) def forward(self, x): seq = x.view(x.shape[0], -1, 3); h = torch.zeros(x.shape[0], self.hdim, device=x.device) self.edges=[]; self.gains=[] for xt in seq.transpose(0,1): z = xt @ self.wih.T + self.bih + (self.gain*h) @ self.whh.T + self.bhh 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)) h = (1-u)*h + u*n if self.controlled: edge=self.estimate_edge(h.detach(), xt.detach()) self.gain *= math.exp(self.alpha*(self.target-self.gain*edge)) self.gain=float(np.clip(self.gain, 0.05, 2.0)) else: edge=self.estimate_edge(h.detach(), xt.detach()) self.edges.append(edge*self.gain); self.gains.append(self.gain) return self.head(h) def train_one(track, seed, controlled, cfg): np.random.seed(seed); torch.manual_seed(seed) ds=bench.get_dataset(track, seed=seed, n_train=400, n_test=200) model=ControlledGRU(ds['out_dim'], controlled, cfg.get('target',.9), cfg.get('alpha',.25)) model, metric, hist=bench.train_model(model, ds, epochs=cfg['epochs'], lr=cfg['lr'], batch=128) 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')} return float(metric), sig def eval_cfg(controlled, cfg, seeds): vals=[]; sig=[] for s in seeds: v,z=train_one('dynamics',s,controlled,cfg); vals.append(v); sig.append(z) return {'per_seed':vals,'mean':float(np.mean(vals)),'signatures':sig} def main(): # Same union of learning rates and equal 3-point search on both sides. seeds=list(range(8)); configs=[{'lr':1e-3,'epochs':12},{'lr':3e-3,'epochs':12},{'lr':6e-3,'epochs':12}] base_sweep=[]; idea_sweep=[] for c in configs: b=eval_cfg(False,c,seeds[:4]); base_sweep.append({'cfg':c,'mean':b['mean']}) i=eval_cfg(True,c,seeds[:4]); idea_sweep.append({'cfg':c,'mean':i['mean']}) bc=min(configs,key=lambda c: next(x['mean'] for x in base_sweep if x['cfg']==c)) ic=min(configs,key=lambda c: next(x['mean'] for x in idea_sweep if x['cfg']==c)) baseline=eval_cfg(False,bc,seeds); idea=eval_cfg(True,ic,seeds) diffs=[a-b for a,b in zip(idea['per_seed'],baseline['per_seed'])] p=bench.permutation_pvalue(diffs) pred=ic.get('target',.9); observed=float(np.mean([x['edge_mean'] for x in idea['signatures']])) report={'bench_version':1,'track':'dynamics','model':'rnn_small','metric_direction':'lower is better','n_seeds':8, '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']}, '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'}, '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)}} with open('bench_report.json','w') as f: json.dump(report,f,indent=2) print(json.dumps(report,indent=2)) if __name__=='__main__': main()