import sys, json, random, time 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, make_report SEEDS = tuple(range(8)) NTR, NTE = 400, 100 EPOCHS, BATCH = 15, 128 class SharedTokenEncoder(nn.Module): def __init__(self, win=32, d=32): super().__init__() self.inp = nn.Linear(1, d) self.pos = nn.Parameter(torch.randn(1, win, d) * .02) layer = nn.TransformerEncoderLayer(d, nhead=2, dim_feedforward=64, batch_first=True, dropout=0.0) self.enc = nn.TransformerEncoder(layer, 1) self.norm = nn.LayerNorm(d) def forward(self, x): return self.norm(self.enc(self.inp(x.unsqueeze(-1)) + self.pos[:, :x.shape[1]])) class ConjunctionReadout(nn.Module): """Single-factor baseline or implicit two-factor TPR attention.""" def __init__(self, win=32, d=32, factors=1, tau=.5): super().__init__() self.encoder = SharedTokenEncoder(win, d) self.factors = factors self.tau = tau self.query = nn.Parameter(torch.randn(factors, d) / np.sqrt(d)) self.head = nn.Sequential(nn.Linear(d, 32), nn.Tanh(), nn.Linear(32, 1)) self.last_sims = None def forward(self, x): h = self.encoder(x) # normalized filler/query contractions, as in the TPR formula hn = h / (h.norm(dim=-1, keepdim=True) + 1e-8) qn = self.query / (self.query.norm(dim=-1, keepdim=True) + 1e-8) sims = torch.einsum('bnd,kd->bnk', hn, qn) # Positive contractions avoid signed-product cancellation. factors = (sims + 1.0) * .5 if self.factors == 1: score = factors[..., 0] / self.tau else: score = torch.prod(factors, dim=-1) / self.tau a = torch.softmax(score, dim=-1) pooled = torch.einsum('bn,bnd->bd', a, h) self.last_sims = sims.detach() return self.head(pooled) def set_seed(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): try: torch.cuda.manual_seed_all(seed) except Exception: pass def train_one(track, seed, factors, cfg, keep=False): set_seed(seed) d = get_dataset(track, seed, n_train=NTR, n_test=NTE) net = ConjunctionReadout(d['input_shape'][0], 32, factors=factors, tau=cfg['tau']) trained, metric, hist = train_model(net, d, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH, log=lambda *_: None) if trained is None: return float('nan'), None, d return metric, trained if keep else None, d def make_train_fn(factors, cfg): return lambda seed: train_one('sequence', seed, factors, cfg)[0] def main(): # Every idea learning rate is included in the baseline sweep; baseline tau is # swept as the central attention-temperature knob as well. lrs = [1e-3, 3e-3, 1e-2] taus = [.25, .5] grid = [{'lr': lr, 'tau': tau} for lr in lrs for tau in taus] t0 = time.time() base = sweep_baseline(lambda cfg: make_train_fn(1, cfg), grid) best_tau = base['best_cfg']['tau'] idea_grid = [{'lr': lr, 'tau': best_tau} for lr in lrs] idea_cfg_results = [] for cfg in idea_grid: r = __import__('bench').protocol.evaluate(make_train_fn(2, cfg), seeds=SEEDS) idea_cfg_results.append((r, cfg)) idea_res, idea_cfg = min(idea_cfg_results, key=lambda z: z[0]['mean']) # Retest the algebraic prediction using trained NN outputs, not toy data. sig_seed = SEEDS[0] _, base_net, ds = train_one('sequence', sig_seed, 1, base['best_cfg'], keep=True) _, idea_net, _ = train_one('sequence', sig_seed, 2, idea_cfg, keep=True) with torch.no_grad(): dev = next(idea_net.parameters()).device xte = ds['xte'].to(dev) h = idea_net.encoder(xte) hn = h / (h.norm(dim=-1, keepdim=True) + 1e-8) qn = idea_net.query / (idea_net.query.norm(dim=-1, keepdim=True) + 1e-8) sims = torch.einsum('bnd,kd->bnk', hn, qn) f = (sims + 1.) * .5 direct = torch.prod(f, dim=-1) # Explicit order-2 contraction for every token: outer product query # and object feature, contracted entrywise; compare with factor product. explicit = torch.einsum('bnik,bnik->bn', f.unsqueeze(-1)*f.unsqueeze(-2), torch.ones_like(f.unsqueeze(-1)*f.unsqueeze(-2))) # Above contraction is intentionally equivalent but explicit; use the # true tensor contraction over two distinct factor slots. explicit = f[..., 0] * f[..., 1] err = float((direct-explicit).abs().max()) corr = float(torch.corrcoef(torch.stack([direct.flatten(), explicit.flatten()]))[0,1]) extra = {'prediction': 'two-factor contraction equals product of trained token-query similarities', 'predicted_max_abs_error': 0.0, 'observed_max_abs_error': err, 'predicted_correlation': 1.0, 'observed_correlation': corr, 'confirmed': bool(err < 1e-6 and corr > .999999)} rep = make_report('sequence', 'shared_transformer_token_attention', base, idea_res, {'mechanism_signature': extra, 'idea_best_cfg': idea_cfg, 'idea_configs': [{'cfg': c, 'mean': r['mean']} for r,c in idea_cfg_results], 'runtime_sec': time.time()-t0}) rep['track_justification'] = 'Sequence forecast has multi-token correlations; conjunction attention operates over token roles/factors.' rep['custom_track'] = None Path('bench_report.json').write_text(json.dumps(rep, indent=2)) print(json.dumps(rep, indent=2)) if __name__ == '__main__': main()