import json import sys from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, train_model, sweep_baseline, evaluate, make_report SEED = 123 EPOCHS = 16 BATCH = 128 NTRAIN, NTEST = 1200, 400 K = 12 class OrderedSequenceNet(nn.Module): """Shared sequence message-passing network; beta=0 is symmetric diffusion.""" def __init__(self, win, hidden=48, beta=0.0, eps=0.0): super().__init__() self.win, self.hidden, self.beta = win, hidden, float(beta) # learned scalar ordering from each observed token, as specified self.order = nn.Sequential(nn.Linear(1, 16), nn.Tanh(), nn.Linear(16, 1)) self.value = nn.Linear(1, hidden) self.update = nn.Sequential(nn.Linear(2 * hidden, hidden), nn.Tanh()) self.head = nn.Sequential(nn.Linear(win * hidden, 48), nn.Tanh(), nn.Linear(48, 1)) # Fixed temporal coordinates and kNN graph. eps is estimated from the graph. t = torch.arange(win, dtype=torch.float32) dist = (t[:, None] - t[None, :]).abs() knn = torch.argsort(dist, dim=1)[:, :K] self.register_buffer('nbr', knn) dd = torch.gather(dist ** 2, 1, knn) self.eps = float(eps if eps > 0 else torch.median(dd[:, -1]) / 4.0) self.register_buffer('d2', dd) def weights(self, x): # x [B,W], weights [B,W,K]; measured NN behavior is used in signature. s = self.order(x.unsqueeze(-1)).squeeze(-1) sj = s[:, self.nbr] logits = -self.d2[None, :, :] / (4 * self.eps) + self.beta * (sj - s[:, :, None]) return torch.softmax(logits.clamp(-20, 20), dim=-1) def forward(self, x): h = self.value(x.unsqueeze(-1)) p = self.weights(x) neigh = h[:, self.nbr, :] # B,W,K,H msg = (p.unsqueeze(-1) * neigh).sum(2) h = h + self.update(torch.cat([h, msg], dim=-1)) return self.head(h.reshape(x.shape[0], -1)) def make_fn(cfg): def run(seed): torch.manual_seed(seed); np.random.seed(seed) ds = get_dataset('sequence', seed, n_train=NTRAIN, n_test=NTEST) model = OrderedSequenceNet(ds['input_shape'][0], beta=cfg['beta']) _, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH, log=lambda *_: None) return metric return run def mechanism_signature(seed, cfg): torch.manual_seed(seed); np.random.seed(seed) ds = get_dataset('sequence', seed, n_train=NTRAIN, n_test=NTEST) model = OrderedSequenceNet(ds['input_shape'][0], beta=cfg['beta']) model, _, _ = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH, log=lambda *_: None) model.eval() device = next(model.parameters()).device x = ds['xte'][:128].to(device) with torch.no_grad(): p = model.weights(x) t = torch.arange(model.win, dtype=torch.float32, device=x.device) tj = t[model.nbr] observed = float((p * (tj[None] - t[None, :, None])).sum(-1).mean() / model.eps) s = model.order(x.unsqueeze(-1)).squeeze(-1) grad = [] # predicted local drift uses observed learned scalar finite differences on graph dsj = s[:, model.nbr] - s[:, :, None] pred = float((p * (2 * model.beta * dsj / (tj[None] - t[None,:,None]).clamp_min(1e-6))).nanmean()) if model.beta else 0.0 entropy = float((-p * (p.clamp_min(1e-12).log())).sum(-1).mean()) # For temporal coordinates, compare actual displacement/eps to the direct # NN-scale finite-difference prediction 2 beta * local ds/dt. with torch.no_grad(): dt = (tj[None] - t[None,:,None]) local = torch.where(dt.abs() > 0, dsj / dt, torch.zeros_like(dt)) predicted = float((p * (2 * model.beta * model.eps * local)).sum(-1).mean() / model.eps) if model.beta else 0.0 return {'beta': cfg['beta'], 'epsilon': model.eps, 'observed_drift_over_eps': observed, 'predicted_drift_over_eps': predicted, 'absolute_error': abs(observed-predicted), 'entropy': entropy, 'confirmed': bool(abs(observed-predicted) < 0.35)} def main(): # Union of learning rates is shared by baseline and idea. Baseline's method knob # is explicitly beta=0; idea sweeps the order strength at the same three rates. lrs = [1e-3, 3e-3, 6e-3] base_grid = [{'lr': lr, 'beta': 0.0} for lr in lrs] idea_grid = [{'lr': lr, 'beta': beta} for lr, beta in zip(lrs, [0.75, 1.0, 1.25])] base = sweep_baseline(make_fn, base_grid) # Comparable 3-config idea sweep on the four tuning seeds, then full 8 paired seeds. idea_trials = [{'cfg': c, 'mean': evaluate(make_fn(c), seeds=(0,1,2,3))['mean']} for c in idea_grid] best = min(idea_trials, key=lambda z: z['mean'])['cfg'] idea = evaluate(make_fn(best)) extra = {'mechanism_signature': mechanism_signature(0, best), 'track_choice': 'sequence: the mechanism is local message passing over multi-token temporal correlations, not a single-token task.', 'idea_sweep': idea_trials, 'selected_cfg': best} rep = make_report('sequence', 'ordered_diffusion_sequence', base, idea, extra) Path('bench_report.json').write_text(json.dumps(rep, indent=2)) print(json.dumps(rep, indent=2)) if __name__ == '__main__': main()