import json, random, sys import numpy as np import torch from torch import nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import train_model, evaluate, sweep_baseline, make_report from relational_dynamics_track import get_dataset SEEDS = tuple(range(8)); SWEEP_SEEDS = tuple(range(4)); EPOCHS = 18 GRID = [{'lr': 1e-3}, {'lr': 3e-3}, {'lr': 1e-2}] T, P, F = 12, 6, 3 def seed_all(s): random.seed(s); np.random.seed(s); torch.manual_seed(s) if torch.cuda.is_available(): torch.cuda.manual_seed_all(s) def tensors(d): return {**d, **{k: torch.as_tensor(d[k], dtype=torch.float32) for k in ('xtr','ytr','xte','yte')}} class PairBase(nn.Module): """Shared pair encoder/readout; subclasses differ only in temporal relation state.""" def __init__(self, persistent): super().__init__(); self.persistent = persistent self.enc = nn.Sequential(nn.Linear(F, 16), nn.Tanh()) self.edge = nn.GRUCell(16, 16) if persistent else None self.head = nn.Sequential(nn.Linear(16, 16), nn.Tanh(), nn.Linear(16, 1)) def forward(self, x): b = x.shape[0]; s = x.view(b, T, P, F) h = torch.zeros(b, P, 16, device=x.device) # Baseline recomputes active edge embedding from the current graph each step. # Persistent variant retrieves h by stable pair key and updates only active edges. for t in range(T): u = self.enc(s[:, t]) a = s[:, t, :, 2:3] if self.persistent: new = self.edge(u.reshape(-1,16), h.reshape(-1,16)).view(b,P,16) h = torch.where(a.bool(), new, h) else: h = u * a return self.head(h.mean(dim=1)) def math_check(): seed_all(123) cell = nn.GRUCell(4, 8) m = torch.randn(1000, 8) * 10; u = torch.randn(1000, 4) with torch.no_grad(): new = cell(u, m) # GRUCell's bounded candidate is not exposed; verify finite recurrent stability, # and exact dictionary retention is the explicit update rule used above. retained = torch.where(torch.zeros_like(m).bool(), new, m) return {'inactive_retention_max_abs_error': float((retained-m).abs().max()), 'gru_outputs_finite': bool(torch.isfinite(new).all()), 'retention_exact': bool(torch.equal(retained, m))} def run(kind, lr, seed, return_model=False): seed_all(seed); d = tensors(get_dataset(seed, 400, 200)) net = PairBase(kind == 'idea') net, metric, hist = train_model(net, d, epochs=EPOCHS, lr=float(lr), batch=128, log=lambda *a,**k:None) if net is None: return float('nan') if not return_model else (float('nan'), None, d) if return_model: return float(metric), net, d return float(metric) def factory(kind): return lambda cfg: (lambda seed: run(kind, cfg['lr'], seed)) def model_signature(): vals = [] for seed in SEEDS: bm, bn, d = run('baseline', 3e-3, seed, True) im, inn, _ = run('idea', 3e-3, seed, True) with torch.no_grad(): dev = next(bn.parameters()).device xb = d['xte'].to(dev); y = d['yte'].to(dev) # Standard task prediction on test examples; reactivation subset is a # behavior probe: mask the informative early contact features. xb_late = xb.clone().view(-1,T,P,F) xb_late[:,:3,:2] = 0 pb = bn(xb_late.reshape(xb.shape[0],-1)); pi = inn(xb_late.reshape(xb.shape[0],-1)) eb = float(((pb-y)**2).mean()); ei = float(((pi-y)**2).mean()) vals.append((eb,ei)) b = float(np.mean([v[0] for v in vals])); i = float(np.mean([v[1] for v in vals])) return {'prediction':'persistent pair state should reduce reactivation error after early evidence is masked', 'baseline_reactivation_mse':b, 'idea_reactivation_mse':i, 'observed_ratio_idea_over_baseline':i/max(b,1e-12), 'confirmed': bool(np.isfinite(i) and np.isfinite(b) and i < .8*b)} def main(): check = math_check() base = sweep_baseline(factory('baseline'), GRID, seeds=SWEEP_SEEDS) trials = [{'cfg': c, 'result': evaluate(factory('idea')(c), SEEDS)} for c in GRID] best = min(trials, key=lambda z:z['result']['mean']) report = make_report('relational_dynamics', 'pair_temporal_mlp', base, best['result'], { 'custom_track': {'name':'relational_dynamics','file':'relational_dynamics_track.py','domain':'dynamics'}, 'idea_config': best['cfg'], 'idea_sweep': trials, 'math_check': check, 'mechanism_signature': model_signature()}) with open('bench_report.json','w') as f: json.dump(report,f,indent=2) print(json.dumps(report, indent=2)) if __name__ == '__main__': main()