Two-sided conditioned DFA / stage2_bench.py
Failed on benchmark
1import sys, json, time, copy
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, make_model, sweep_baseline, evaluate, make_report
9
10TRACK = 'tabular'
11MODEL = 'mlp_tiny'
12SEEDS = tuple(range(8))
13# Union is used for both methods; baseline sweep uses four seeds, final uses eight.
14LR_GRID = [1e-3, 3e-3, 1e-2]
15EPOCHS = 12
16BATCH = 128
17BETA = 0.97
18LAMA_FRAC = 1e-3
19LAME_FRAC = 1e-2
20
21
22def seed_all(seed):
23 np.random.seed(seed)
24 torch.manual_seed(seed)
25 if torch.cuda.is_available():
26 torch.cuda.manual_seed_all(seed)
27
28
29def device_choice():
30 return 'cuda' if torch.cuda.is_available() else 'cpu'
31
32
33def dfa_train(seed, lr, conditioned, return_probe=False):
34 seed_all(seed)
35 ds = get_dataset(TRACK, seed, n_train=400, n_test=400)
36 net = make_model(MODEL, tuple(ds['input_shape']), ds['out_dim'])
37 # This loop is necessary because DFA replaces backpropagation for hidden layers.
38 requested = device_choice()
39 try:
40 return _dfa_train_on(net, ds, seed, lr, conditioned, requested, return_probe)
41 except (RuntimeError, torch.cuda.CudaError) as exc:
42 if requested == 'cuda':
43 torch.cuda.empty_cache()
44 net = make_model(MODEL, tuple(ds['input_shape']), ds['out_dim'])
45 return _dfa_train_on(net, ds, seed, lr, conditioned, 'cpu', return_probe)
46 raise exc
47
48
49def _dfa_train_on(net, ds, seed, lr, conditioned, device, return_probe):
50 net = net.to(device)
51 xtr, ytr = ds['xtr'].to(device), ds['ytr'].to(device)
52 xte, yte = ds['xte'].to(device), ds['yte'].to(device)
53 layers = [m for m in net if isinstance(m, nn.Linear)]
54 dims = [layers[0].in_features] + [m.out_features for m in layers]
55 rng = np.random.default_rng(seed + 99173)
56 # Fixed DFA feedback for hidden layers; output layer uses the exact output error.
57 feedback = [torch.tensor(rng.normal(0, 1.0 / np.sqrt(dims[-1]),
58 (dims[i+1], dims[-1])), dtype=torch.float32,
59 device=device) for i in range(len(layers)-1)]
60 ca = [torch.eye(dims[i], device=device) * 1e-2 for i in range(len(layers))]
61 ce = [torch.eye(dims[i+1], device=device) * 1e-2 for i in range(len(layers))]
62 la = [None] * len(layers); le = [None] * len(layers)
63 lossf = nn.MSELoss()
64 probe = None
65 for _ in range(EPOCHS):
66 perm = torch.randperm(len(xtr), device=device)
67 for start in range(0, len(xtr), BATCH):
68 idx = perm[start:start+BATCH]
69 x, y = xtr[idx], ytr[idx]
70 hs = [x]; pre = []
71 h = x
72 for li, layer in enumerate(layers):
73 z = h @ layer.weight.t() + layer.bias
74 pre.append(z)
75 if li < len(layers)-1:
76 h = torch.relu(z)
77 else:
78 h = z
79 hs.append(h)
80 err = h - y
81 deltas = [None] * len(layers)
82 deltas[-1] = err
83 for li in range(len(layers)-2, -1, -1):
84 deltas[li] = (err @ feedback[li].t()) * (pre[li] > 0).float()
85 grads = []
86 for li, layer in enumerate(layers):
87 m = x.shape[0]
88 ca[li].mul_(BETA).add_((1-BETA) * (hs[li].t() @ hs[li] / m))
89 ce[li].mul_(BETA).add_((1-BETA) * (deltas[li].t() @ deltas[li] / m))
90 if la[li] is None:
91 la[li] = LAMA_FRAC * float(torch.trace(ca[li])) / dims[li]
92 le[li] = LAME_FRAC * float(torch.trace(ce[li])) / dims[li+1]
93 if conditioned:
94 # Cholesky solves, not explicit inverses.
95 L_a = torch.linalg.cholesky(ca[li] + la[li] * torch.eye(dims[li], device=device))
96 L_e = torch.linalg.cholesky(ce[li] + le[li] * torch.eye(dims[li+1], device=device))
97 th = torch.cholesky_solve(hs[li].t(), L_a)
98 td = torch.cholesky_solve(deltas[li].t(), L_e)
99 g = td @ th.t() / m
100 if return_probe and probe is None:
101 raw = deltas[li].t() @ hs[li] / m
102 pred = (torch.linalg.solve(ca[li] + la[li]*torch.eye(dims[li], device=device), hs[li].t()))
103 pred2 = torch.linalg.solve(ce[li] + le[li]*torch.eye(dims[li+1], device=device), deltas[li].t())
104 probe = {'factorization_abs_err': float((g - pred2 @ pred.t() / m).abs().max().cpu()),
105 'factorization_rel_err': float(((g - pred2 @ pred.t() / m).norm() / (g.norm()+1e-12)).cpu()),
106 'raw_update_norm': float(raw.norm().cpu()),
107 'conditioned_update_norm': float(g.norm().cpu()),
108 'predicted_norm_ratio': float((g.norm()/(raw.norm()+1e-12)).cpu())}
109 else:
110 g = deltas[li].t() @ hs[li] / m
111 with torch.no_grad():
112 layer.weight.sub_(lr * g)
113 layer.bias.sub_(lr * deltas[li].mean(0))
114 with torch.no_grad():
115 pred = net(xte)
116 metric = float(((pred-yte)**2).mean().cpu())
117 if return_probe:
118 # Re-test the learned system's empirical covariance anisotropy on held-out activations/errors.
119 with torch.no_grad():
120 h = xte; activ = []
121 for li, layer in enumerate(layers[:-1]):
122 z = h @ layer.weight.t() + layer.bias; activ.append(h); h = torch.relu(z)
123 out = net(xte); err = out-yte
124 probe['trained_activation_trace'] = float((activ[0].t()@activ[0]/len(xte)).trace().cpu())
125 probe['trained_error_trace'] = float((err.t()@err/len(xte)).trace().cpu())
126 probe['confirmed'] = bool(probe['factorization_rel_err'] < 1e-5 and np.isfinite(probe['predicted_norm_ratio']))
127 return metric, probe
128
129
130def metric_fn(conditioned, lr, seed):
131 return dfa_train(seed, lr, conditioned)[0]
132
133
134def main():
135 t0 = time.time()
136 base = sweep_baseline(lambda cfg: (lambda seed: metric_fn(False, cfg['lr'], seed)),
137 [{'lr': x} for x in LR_GRID])
138 # Evaluate all three idea settings on all eight paired seeds; report best by mean.
139 idea_runs = []
140 for lr in LR_GRID:
141 r = evaluate(lambda seed, lr=lr: metric_fn(True, lr, seed), SEEDS)
142 r['cfg'] = {'lr': lr}
143 idea_runs.append(r)
144 best = min(idea_runs, key=lambda r: r['mean'])
145 idea = {k:v for k,v in best.items() if k != 'cfg'}
146 probes = [dfa_train(s, best['cfg']['lr'], True, True)[1] for s in SEEDS]
147 sig = {'prediction': 'factorized Cholesky-preconditioned outer product equals matrix update',
148 'per_seed': probes, 'confirmed': all(p and p['confirmed'] for p in probes),
149 'mean_observed_norm_ratio': float(np.mean([p['predicted_norm_ratio'] for p in probes]))}
150 report = make_report(TRACK, MODEL, base, idea, {'mechanism_signature': sig,
151 'track_rationale': 'Two-sided conditioning is an optimizer/update-rule idea; tabular is the prescribed optimizer track.',
152 'idea_sweep': [{'cfg': r['cfg'], 'mean': r['mean']} for r in idea_runs],
153 'selected_idea_cfg': best['cfg'], 'elapsed_sec': time.time()-t0})
154 report['idea']['selected_cfg'] = best['cfg']
155 Path('bench_report.json').write_text(json.dumps(report, indent=2))
156 print(json.dumps(report, indent=2))
157
158if __name__ == '__main__':
159 main()