Harmonic-Mode Branch for Topological Memory / bench_harmonic.py
Mechanism confirmed, baseline not beaten
1import sys, json, random
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
8
9SEEDS = tuple(range(8))
10DIM, K = 64, 2
11TRACK = 'cycle_topology_graph'
12MODEL = 'mlp_tiny'
13
14
15def seed_all(seed):
16 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
17 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
18
19
20class TopologicalMLP(nn.Module):
21 """Same network in both arms; only latent transition dissipation differs."""
22 def __init__(self, input_dim, out_dim, mode, damping):
23 super().__init__()
24 self.mode, self.damping = mode, damping
25 self.enc = nn.Sequential(nn.Linear(input_dim, DIM), nn.ReLU(), nn.Linear(DIM, DIM))
26 self.interaction = nn.Sequential(nn.Tanh(), nn.Linear(DIM, DIM), nn.Tanh())
27 self.head = nn.Linear(DIM, out_dim)
28 H = torch.zeros(DIM, K); H[0, 0] = 1.; H[1, 1] = 1.
29 self.register_buffer('H', H)
30 self.register_buffer('PH', H @ H.T)
31 self.register_buffer('PP', torch.eye(DIM) - H @ H.T)
32
33 def forward(self, x):
34 z = self.enc(x)
35 n = self.interaction(z)
36 if self.mode == 'baseline':
37 znext = z + n - self.damping * z
38 else:
39 h = z @ self.PH
40 u = z @ self.PP
41 # Projected interaction and damping preserve the harmonic branch.
42 nperp = n @ self.PP
43 znext = h + u + nperp - self.damping * u
44 return self.head(znext)
45
46 def signature(self, x):
47 with torch.no_grad():
48 x = x.to(next(self.parameters()).device)
49 z = self.enc(x); n = self.interaction(z)
50 h = z @ self.PH; u = z @ self.PP
51 if self.mode == 'baseline':
52 step = n - self.damping * z
53 else:
54 step = n @ self.PP - self.damping * u
55 harmonic_step = step @ self.PH
56 diss_step = step @ self.PP
57 return float(harmonic_step.norm(dim=1).mean()), float(diss_step.norm(dim=1).mean()), float(h.norm(dim=1).mean())
58
59
60def train_one(mode, cfg, seed, return_model=False):
61 seed_all(seed)
62 d = get_dataset(TRACK, seed, n_train=400, n_test=400)
63 net = TopologicalMLP(int(np.prod(d['input_shape'])), d['out_dim'], mode, cfg['damping'])
64 net, metric, hist = train_model(net, d, epochs=cfg['epochs'], lr=cfg['lr'], batch=128)
65 if return_model:
66 return metric, net, d
67 return metric
68
69
70def main():
71 # Union parity: every lr and damping value is tried by baseline and idea.
72 grid = [{'lr': lr, 'damping': damp, 'epochs': 18}
73 for lr in (1e-3, 3e-3, 1e-2) for damp in (0.05, 0.2, 0.5)]
74 base = sweep_baseline(lambda c: lambda s: train_one('baseline', c, s), grid, seeds=(0,1,2,3))
75 idea_cfgs = grid
76 idea_trials = []
77 for cfg in idea_cfgs:
78 r = evaluate(lambda s, c=cfg: train_one('idea', c, s), seeds=(0,1,2,3))
79 idea_trials.append({'cfg': cfg, 'mean': r['mean']})
80 best = min(idea_trials, key=lambda q: q['mean'])['cfg']
81 idea = evaluate(lambda s: train_one('idea', best, s), seeds=SEEDS)
82 # Trained-model signature: measured projected harmonic update and dissipative update.
83 hs, ds, hm = [], [], []
84 for s in SEEDS:
85 _, model, data = train_one('idea', best, s, return_model=True)
86 h, u, a = model.signature(data['xte']); hs.append(h); ds.append(u); hm.append(a)
87 base['idea_grid'] = idea_trials
88 base['baseline_grid_union'] = grid
89 rep = make_report(TRACK, MODEL, base, idea, extra={
90 'prediction': 'harmonic update is near zero while dissipative update remains nonzero',
91 'trained_model_observed_mean_harmonic_step': float(np.mean(hs)),
92 'trained_model_observed_mean_dissipative_step': float(np.mean(ds)),
93 'trained_model_observed_mean_harmonic_amplitude': float(np.mean(hm)),
94 'confirmed': bool(np.mean(hs) < 1e-7 and np.mean(ds) > 1e-5),
95 'measurement': 'mean latent transition norms on each trained idea model test split'
96 })
97 rep['custom_track'] = {'name': TRACK, 'file': '/home/maxwelhelp/all/math2nn/bench/custom_tracks/cycle_topology_graph.py', 'domain': 'graph_topology'}
98 Path('bench_report.json').write_text(json.dumps(rep, indent=2))
99 print(json.dumps(rep, indent=2))
100
101if __name__ == '__main__':
102 main()