Two-sided conditioned DFA / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
  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()