Ordered Diffusion Message Passing / ordered_diffusion_bench.py
Mechanism confirmed, baseline not beaten
1import json
2import sys
3from pathlib import Path
4import numpy as np
5import torch
6import torch.nn as nn
7
8sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
9from bench import get_dataset, train_model, sweep_baseline, evaluate, make_report
10
11SEED = 123
12EPOCHS = 16
13BATCH = 128
14NTRAIN, NTEST = 1200, 400
15K = 12
16
17class OrderedSequenceNet(nn.Module):
18 """Shared sequence message-passing network; beta=0 is symmetric diffusion."""
19 def __init__(self, win, hidden=48, beta=0.0, eps=0.0):
20 super().__init__()
21 self.win, self.hidden, self.beta = win, hidden, float(beta)
22 # learned scalar ordering from each observed token, as specified
23 self.order = nn.Sequential(nn.Linear(1, 16), nn.Tanh(), nn.Linear(16, 1))
24 self.value = nn.Linear(1, hidden)
25 self.update = nn.Sequential(nn.Linear(2 * hidden, hidden), nn.Tanh())
26 self.head = nn.Sequential(nn.Linear(win * hidden, 48), nn.Tanh(), nn.Linear(48, 1))
27 # Fixed temporal coordinates and kNN graph. eps is estimated from the graph.
28 t = torch.arange(win, dtype=torch.float32)
29 dist = (t[:, None] - t[None, :]).abs()
30 knn = torch.argsort(dist, dim=1)[:, :K]
31 self.register_buffer('nbr', knn)
32 dd = torch.gather(dist ** 2, 1, knn)
33 self.eps = float(eps if eps > 0 else torch.median(dd[:, -1]) / 4.0)
34 self.register_buffer('d2', dd)
35
36 def weights(self, x):
37 # x [B,W], weights [B,W,K]; measured NN behavior is used in signature.
38 s = self.order(x.unsqueeze(-1)).squeeze(-1)
39 sj = s[:, self.nbr]
40 logits = -self.d2[None, :, :] / (4 * self.eps) + self.beta * (sj - s[:, :, None])
41 return torch.softmax(logits.clamp(-20, 20), dim=-1)
42
43 def forward(self, x):
44 h = self.value(x.unsqueeze(-1))
45 p = self.weights(x)
46 neigh = h[:, self.nbr, :] # B,W,K,H
47 msg = (p.unsqueeze(-1) * neigh).sum(2)
48 h = h + self.update(torch.cat([h, msg], dim=-1))
49 return self.head(h.reshape(x.shape[0], -1))
50
51def make_fn(cfg):
52 def run(seed):
53 torch.manual_seed(seed); np.random.seed(seed)
54 ds = get_dataset('sequence', seed, n_train=NTRAIN, n_test=NTEST)
55 model = OrderedSequenceNet(ds['input_shape'][0], beta=cfg['beta'])
56 _, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH, log=lambda *_: None)
57 return metric
58 return run
59
60def mechanism_signature(seed, cfg):
61 torch.manual_seed(seed); np.random.seed(seed)
62 ds = get_dataset('sequence', seed, n_train=NTRAIN, n_test=NTEST)
63 model = OrderedSequenceNet(ds['input_shape'][0], beta=cfg['beta'])
64 model, _, _ = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH, log=lambda *_: None)
65 model.eval()
66 device = next(model.parameters()).device
67 x = ds['xte'][:128].to(device)
68 with torch.no_grad():
69 p = model.weights(x)
70 t = torch.arange(model.win, dtype=torch.float32, device=x.device)
71 tj = t[model.nbr]
72 observed = float((p * (tj[None] - t[None, :, None])).sum(-1).mean() / model.eps)
73 s = model.order(x.unsqueeze(-1)).squeeze(-1)
74 grad = []
75 # predicted local drift uses observed learned scalar finite differences on graph
76 dsj = s[:, model.nbr] - s[:, :, None]
77 pred = float((p * (2 * model.beta * dsj / (tj[None] - t[None,:,None]).clamp_min(1e-6))).nanmean()) if model.beta else 0.0
78 entropy = float((-p * (p.clamp_min(1e-12).log())).sum(-1).mean())
79 # For temporal coordinates, compare actual displacement/eps to the direct
80 # NN-scale finite-difference prediction 2 beta * local ds/dt.
81 with torch.no_grad():
82 dt = (tj[None] - t[None,:,None])
83 local = torch.where(dt.abs() > 0, dsj / dt, torch.zeros_like(dt))
84 predicted = float((p * (2 * model.beta * model.eps * local)).sum(-1).mean() / model.eps) if model.beta else 0.0
85 return {'beta': cfg['beta'], 'epsilon': model.eps, 'observed_drift_over_eps': observed,
86 'predicted_drift_over_eps': predicted, 'absolute_error': abs(observed-predicted),
87 'entropy': entropy, 'confirmed': bool(abs(observed-predicted) < 0.35)}
88
89def main():
90 # Union of learning rates is shared by baseline and idea. Baseline's method knob
91 # is explicitly beta=0; idea sweeps the order strength at the same three rates.
92 lrs = [1e-3, 3e-3, 6e-3]
93 base_grid = [{'lr': lr, 'beta': 0.0} for lr in lrs]
94 idea_grid = [{'lr': lr, 'beta': beta} for lr, beta in zip(lrs, [0.75, 1.0, 1.25])]
95 base = sweep_baseline(make_fn, base_grid)
96 # Comparable 3-config idea sweep on the four tuning seeds, then full 8 paired seeds.
97 idea_trials = [{'cfg': c, 'mean': evaluate(make_fn(c), seeds=(0,1,2,3))['mean']} for c in idea_grid]
98 best = min(idea_trials, key=lambda z: z['mean'])['cfg']
99 idea = evaluate(make_fn(best))
100 extra = {'mechanism_signature': mechanism_signature(0, best),
101 'track_choice': 'sequence: the mechanism is local message passing over multi-token temporal correlations, not a single-token task.',
102 'idea_sweep': idea_trials, 'selected_cfg': best}
103 rep = make_report('sequence', 'ordered_diffusion_sequence', base, idea, extra)
104 Path('bench_report.json').write_text(json.dumps(rep, indent=2))
105 print(json.dumps(rep, indent=2))
106
107if __name__ == '__main__':
108 main()