Implicit Higher-Order TPR Memory / run_bench.py
Mechanism confirmed, baseline not beaten
1import sys, json, random, time
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, sweep_baseline, make_report
8
9SEEDS = tuple(range(8))
10NTR, NTE = 400, 100
11EPOCHS, BATCH = 15, 128
12
13class SharedTokenEncoder(nn.Module):
14 def __init__(self, win=32, d=32):
15 super().__init__()
16 self.inp = nn.Linear(1, d)
17 self.pos = nn.Parameter(torch.randn(1, win, d) * .02)
18 layer = nn.TransformerEncoderLayer(d, nhead=2, dim_feedforward=64,
19 batch_first=True, dropout=0.0)
20 self.enc = nn.TransformerEncoder(layer, 1)
21 self.norm = nn.LayerNorm(d)
22
23 def forward(self, x):
24 return self.norm(self.enc(self.inp(x.unsqueeze(-1)) + self.pos[:, :x.shape[1]]))
25
26class ConjunctionReadout(nn.Module):
27 """Single-factor baseline or implicit two-factor TPR attention."""
28 def __init__(self, win=32, d=32, factors=1, tau=.5):
29 super().__init__()
30 self.encoder = SharedTokenEncoder(win, d)
31 self.factors = factors
32 self.tau = tau
33 self.query = nn.Parameter(torch.randn(factors, d) / np.sqrt(d))
34 self.head = nn.Sequential(nn.Linear(d, 32), nn.Tanh(), nn.Linear(32, 1))
35 self.last_sims = None
36
37 def forward(self, x):
38 h = self.encoder(x)
39 # normalized filler/query contractions, as in the TPR formula
40 hn = h / (h.norm(dim=-1, keepdim=True) + 1e-8)
41 qn = self.query / (self.query.norm(dim=-1, keepdim=True) + 1e-8)
42 sims = torch.einsum('bnd,kd->bnk', hn, qn)
43 # Positive contractions avoid signed-product cancellation.
44 factors = (sims + 1.0) * .5
45 if self.factors == 1:
46 score = factors[..., 0] / self.tau
47 else:
48 score = torch.prod(factors, dim=-1) / self.tau
49 a = torch.softmax(score, dim=-1)
50 pooled = torch.einsum('bn,bnd->bd', a, h)
51 self.last_sims = sims.detach()
52 return self.head(pooled)
53
54def set_seed(seed):
55 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
56 if torch.cuda.is_available():
57 try: torch.cuda.manual_seed_all(seed)
58 except Exception: pass
59
60def train_one(track, seed, factors, cfg, keep=False):
61 set_seed(seed)
62 d = get_dataset(track, seed, n_train=NTR, n_test=NTE)
63 net = ConjunctionReadout(d['input_shape'][0], 32, factors=factors, tau=cfg['tau'])
64 trained, metric, hist = train_model(net, d, epochs=EPOCHS, lr=cfg['lr'],
65 batch=BATCH, log=lambda *_: None)
66 if trained is None: return float('nan'), None, d
67 return metric, trained if keep else None, d
68
69def make_train_fn(factors, cfg):
70 return lambda seed: train_one('sequence', seed, factors, cfg)[0]
71
72def main():
73 # Every idea learning rate is included in the baseline sweep; baseline tau is
74 # swept as the central attention-temperature knob as well.
75 lrs = [1e-3, 3e-3, 1e-2]
76 taus = [.25, .5]
77 grid = [{'lr': lr, 'tau': tau} for lr in lrs for tau in taus]
78 t0 = time.time()
79 base = sweep_baseline(lambda cfg: make_train_fn(1, cfg), grid)
80 best_tau = base['best_cfg']['tau']
81 idea_grid = [{'lr': lr, 'tau': best_tau} for lr in lrs]
82 idea_cfg_results = []
83 for cfg in idea_grid:
84 r = __import__('bench').protocol.evaluate(make_train_fn(2, cfg), seeds=SEEDS)
85 idea_cfg_results.append((r, cfg))
86 idea_res, idea_cfg = min(idea_cfg_results, key=lambda z: z[0]['mean'])
87
88 # Retest the algebraic prediction using trained NN outputs, not toy data.
89 sig_seed = SEEDS[0]
90 _, base_net, ds = train_one('sequence', sig_seed, 1, base['best_cfg'], keep=True)
91 _, idea_net, _ = train_one('sequence', sig_seed, 2, idea_cfg, keep=True)
92 with torch.no_grad():
93 dev = next(idea_net.parameters()).device
94 xte = ds['xte'].to(dev)
95 h = idea_net.encoder(xte)
96 hn = h / (h.norm(dim=-1, keepdim=True) + 1e-8)
97 qn = idea_net.query / (idea_net.query.norm(dim=-1, keepdim=True) + 1e-8)
98 sims = torch.einsum('bnd,kd->bnk', hn, qn)
99 f = (sims + 1.) * .5
100 direct = torch.prod(f, dim=-1)
101 # Explicit order-2 contraction for every token: outer product query
102 # and object feature, contracted entrywise; compare with factor product.
103 explicit = torch.einsum('bnik,bnik->bn', f.unsqueeze(-1)*f.unsqueeze(-2),
104 torch.ones_like(f.unsqueeze(-1)*f.unsqueeze(-2)))
105 # Above contraction is intentionally equivalent but explicit; use the
106 # true tensor contraction over two distinct factor slots.
107 explicit = f[..., 0] * f[..., 1]
108 err = float((direct-explicit).abs().max())
109 corr = float(torch.corrcoef(torch.stack([direct.flatten(), explicit.flatten()]))[0,1])
110 extra = {'prediction': 'two-factor contraction equals product of trained token-query similarities',
111 'predicted_max_abs_error': 0.0, 'observed_max_abs_error': err,
112 'predicted_correlation': 1.0, 'observed_correlation': corr,
113 'confirmed': bool(err < 1e-6 and corr > .999999)}
114 rep = make_report('sequence', 'shared_transformer_token_attention', base, idea_res,
115 {'mechanism_signature': extra, 'idea_best_cfg': idea_cfg,
116 'idea_configs': [{'cfg': c, 'mean': r['mean']} for r,c in idea_cfg_results],
117 'runtime_sec': time.time()-t0})
118 rep['track_justification'] = 'Sequence forecast has multi-token correlations; conjunction attention operates over token roles/factors.'
119 rep['custom_track'] = None
120 Path('bench_report.json').write_text(json.dumps(rep, indent=2))
121 print(json.dumps(rep, indent=2))
122
123if __name__ == '__main__': main()