Bifurcation-Aware Adaptive Compute Controller / bench_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, json, random, math, time
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6
  7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  8from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report, permutation_pvalue
  9
 10SEEDS = tuple(range(8))
 11SWEEP_SEEDS = tuple(range(4))
 12EPOCHS = 8
 13BATCH = 128
 14
 15class ControlledRNN(nn.Module):
 16    """Shared GRUCell architecture; baseline and idea differ only in update count."""
 17    def __init__(self, hidden=48, mode='baseline', inner_steps=1, mu0=.03):
 18        super().__init__()
 19        self.inp = nn.Linear(3, hidden)
 20        self.cell = nn.GRUCell(hidden, hidden)
 21        self.head = nn.Linear(hidden, 1)
 22        self.mode = mode
 23        self.inner_steps = int(inner_steps)
 24        self.mu0 = float(mu0)
 25        self.last_stats = {}
 26
 27    def forward(self, x):
 28        seq = x.view(x.shape[0], -1, 3)
 29        h = torch.zeros(x.shape[0], self.cell.hidden_size, device=x.device)
 30        counts, mus = [], []
 31        for token in seq.unbind(1):
 32            q = torch.tanh(self.inp(token))
 33            probe = self.cell(q, h)
 34            z = torch.tanh(h.mean(1))
 35            z_next = torch.tanh(probe.mean(1))
 36            f = z_next - z
 37            mu = f - z.square()
 38            if self.mode == 'baseline':
 39                n = self.inner_steps
 40            else:
 41                # Square-root controller, clipped to safe integer compute levels.
 42                positive = torch.relu(mu.detach())
 43                score = torch.sqrt(positive + 1e-5) / math.sqrt(self.mu0)
 44                n = int(torch.clamp(torch.round(4.0 / (score + 0.15)), 1, 4).max().item())
 45            for _ in range(n):
 46                h = self.cell(q, h)
 47            counts.append(float(n)); mus.append(float(mu.detach().mean()))
 48        self.last_stats = {'mean_updates': float(np.mean(counts)),
 49                           'mean_mu_hat': float(np.mean(mus))}
 50        return self.head(h)
 51
 52def seed_all(seed):
 53    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 54    if torch.cuda.is_available():
 55        torch.cuda.manual_seed_all(seed)
 56
 57def make_train(cfg, mode):
 58    def run(seed):
 59        seed_all(seed)
 60        ds = get_dataset('dynamics', seed, n_train=400, n_test=400)
 61        net = ControlledRNN(mode=mode, inner_steps=cfg['inner_steps'], mu0=cfg['mu0'])
 62        _, metric, _ = train_model(net, ds, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH, log=lambda *_: None)
 63        return float(metric)
 64    return run
 65
 66def main():
 67    # Union parity: every lr and central compute setting is present on both sides.
 68    grid = [
 69        {'lr': 0.0015, 'inner_steps': 1, 'mu0': .02},
 70        {'lr': 0.0030, 'inner_steps': 1, 'mu0': .03},
 71        {'lr': 0.0060, 'inner_steps': 1, 'mu0': .05},
 72    ]
 73    baseline = sweep_baseline(lambda c: make_train(c, 'baseline'), grid, seeds=SWEEP_SEEDS)
 74    idea_sweep = []
 75    for cfg in grid:
 76        r = evaluate(make_train(cfg, 'idea'), seeds=SWEEP_SEEDS)
 77        idea_sweep.append({'cfg': cfg, 'mean': r['mean']})
 78    best = min(idea_sweep, key=lambda x: x['mean'])['cfg']
 79    idea_full = evaluate(make_train(best, 'idea'), seeds=SEEDS)
 80    # Trained-model behavior signature, measured on held-out benchmark examples.
 81    sig_rows = []
 82    for seed in SEEDS:
 83        seed_all(seed)
 84        ds = get_dataset('dynamics', seed, n_train=400, n_test=400)
 85        b = ControlledRNN(mode='baseline', inner_steps=best['inner_steps'], mu0=best['mu0'])
 86        i = ControlledRNN(mode='idea', inner_steps=best['inner_steps'], mu0=best['mu0'])
 87        b, _, _ = train_model(b, ds, epochs=EPOCHS, lr=best['lr'], batch=BATCH, log=lambda *_: None)
 88        i, _, _ = train_model(i, ds, epochs=EPOCHS, lr=best['lr'], batch=BATCH, log=lambda *_: None)
 89        with torch.no_grad():
 90            bdev = next(b.parameters()).device
 91            idev = next(i.parameters()).device
 92            b(ds['xte'].to(bdev)); bs = dict(b.last_stats)
 93            i(ds['xte'].to(idev)); ins = dict(i.last_stats)
 94        sig_rows.append({'seed': seed, 'baseline_updates': bs['mean_updates'],
 95                         'idea_updates': ins['mean_updates'], 'idea_mu_hat': ins['mean_mu_hat']})
 96    observed = float(np.mean([r['idea_updates'] for r in sig_rows]))
 97    predicted = float(np.mean([max(1, min(4, round(4 / (math.sqrt(max(r['idea_mu_hat'],0)+1e-5)/math.sqrt(best['mu0']) + .15)))) for r in sig_rows]))
 98    signature = {'prediction': 'allocation increases as positive mu_hat approaches zero',
 99                 'predicted_mean_updates_from_measured_mu': predicted,
100                 'observed_mean_updates_on_trained_models': observed,
101                 'per_seed': sig_rows, 'confirmed': bool(observed >= 1.0 and predicted >= 1.0 and abs(observed-predicted) <= 1.0)}
102    report = make_report('dynamics', 'rnn_small', baseline, idea_full, signature)
103    report['idea_sweep'] = idea_sweep
104    report['protocol'] = {'paired_seeds': list(SEEDS), 'sweep_seeds': list(SWEEP_SEEDS), 'epochs': EPOCHS, 'batch': BATCH,
105                          'structural_match': 'dynamics/control', 'same_architecture': True}
106    Path('bench_report.json').write_text(json.dumps(report, indent=2))
107    print(json.dumps(report, indent=2))
108
109if __name__ == '__main__':
110    main()