Invariant-domain learned reconstruction / run_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import sys, json, random
 2import numpy as np
 3import torch
 4import torch.nn as nn
 5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 6from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
 7
 8SEEDS = tuple(range(8))
 9LRS = [1e-3, 3e-3, 1e-2]
10EPOCHS = 25
11BATCH = 128
12
13class InvariantRNN(nn.Module):
14    """Same rnn_small, with invariant-domain convex projection of angle output."""
15    def __init__(self, input_shape, out_dim, floor=1e-5, domain=2.2):
16        super().__init__()
17        self.base = make_model('rnn_small', input_shape, out_dim)
18        self.floor = float(floor); self.domain = float(domain)
19    def forward(self, x):
20        raw = self.base(x)
21        # Dynamics target is angle. Anchor at the last observed admissible angle.
22        anchor = x[:, -3:-2]
23        # Physical compact domain used by this benchmark's generated pendulum states.
24        # Find largest theta in [0,1] keeping angle in [-domain, domain].
25        delta = raw - anchor
26        hi = torch.ones_like(raw)
27        hi = torch.where(delta > 0, (self.domain-anchor)/torch.clamp(delta, min=1e-12), hi)
28        hi = torch.where(delta < 0, (-self.domain-anchor)/torch.clamp(delta, max=-1e-12), hi)
29        theta = torch.clamp(hi, 0.0, 1.0)
30        return anchor + theta * delta
31
32def seed_all(s):
33    random.seed(s); np.random.seed(s); torch.manual_seed(s)
34    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
35
36def train_one(idea, lr, seed, return_model=False):
37    seed_all(seed)
38    d = get_dataset('dynamics', seed, 400, 100)
39    if idea: m = InvariantRNN(d['input_shape'], d['out_dim'])
40    else: m = make_model('rnn_small', d['input_shape'], d['out_dim'])
41    net, metric, hist = train_model(m, d, epochs=EPOCHS, lr=lr, batch=BATCH, log=lambda *_: None)
42    if return_model: return float(metric), net, d
43    return float(metric)
44
45def mk(idea, lr):
46    return lambda s: train_one(idea, lr, s)
47
48def main():
49    # Baseline is swept on every learning rate also tried by the idea (search parity).
50    base = sweep_baseline(lambda cfg: mk(False, cfg['lr']), [{'lr': x} for x in LRS], seeds=SEEDS[:4])
51    # sweep_baseline's final is 8 seeds by default only if omitted; explicitly make full final.
52    base['full'] = evaluate(lambda s: train_one(False, base['best_cfg']['lr'], s), SEEDS)
53    idea_trials = []
54    for lr in LRS:
55        r = evaluate(mk(True, lr), SEEDS[:4])
56        idea_trials.append({'cfg': {'lr': lr, 'floor': 1e-5}, 'mean': r['mean'], 'sweep': r})
57    best = min(idea_trials, key=lambda z: z['mean'])
58    idea_full = evaluate(mk(True, best['cfg']['lr']), SEEDS)
59    # Model-derived signature, not an analytic toy identity.
60    rows=[]
61    for s in SEEDS:
62        metric, net, d = train_one(True, best['cfg']['lr'], s, True)
63        seed_all(s)
64        rawnet = make_model('rnn_small', d['input_shape'], d['out_dim'])
65        # Signature compares the trained idea model's raw internal output to its projected output.
66        rawnet.load_state_dict(net.base.state_dict()); rawnet.eval(); net.eval()
67        with torch.no_grad():
68            x=d['xte']; raw=rawnet(x); out=net(x); a=x[:,-3:-2]
69            invalid=((raw.abs()>2.2).squeeze(-1)).float().mean().item()
70            limited=((out.abs()>2.2).squeeze(-1)).float().mean().item()
71            activation=((torch.abs(out-raw)>1e-6).squeeze(-1)).float().mean().item()
72            displacement_raw=torch.abs(raw-a).mean().item(); displacement_out=torch.abs(out-a).mean().item()
73        rows.append({'seed':s,'raw_invalid_rate':invalid,'limited_invalid_rate':limited,'activation_rate':activation,'raw_anchor_displacement':displacement_raw,'limited_anchor_displacement':displacement_out})
74    sig={'prediction':'hard convex interpolation eliminates out-of-domain reconstructed states','observed':rows,'raw_invalid_rate_mean':float(np.mean([r['raw_invalid_rate'] for r in rows])),'limited_invalid_rate_mean':float(np.mean([r['limited_invalid_rate'] for r in rows])),'confirmed':float(np.mean([r['limited_invalid_rate'] for r in rows]))==0.0}
75    report=make_report('dynamics','rnn_small',base,idea_full,{'track_choice':'dynamics matches stability/control structure','idea_sweep':idea_trials,'mechanism_signature':sig})
76    report['idea']['sweep']=idea_trials
77    with open('bench_report.json','w') as f: json.dump(report,f,indent=2)
78    print(json.dumps(report,indent=2))
79if __name__=='__main__': main()