Coupled multilevel gradients for Markov-stream training / markov_stream_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import os, sys, json, math, random
  2import numpy as np
  3import torch
  4import torch.nn as nn
  5
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import make_model, sweep_baseline, make_report
  8from bench.protocol import evaluate, DEFAULT_SEEDS
  9
 10META = {
 11    'name': 'markov_stream_regression',
 12    'domain': 'optimizer',
 13    'description': 'Ordered stationary AR(1) covariate stream with nonlinear regression targets; tests gradient estimation under Markov dependence.'
 14}
 15
 16
 17def get_dataset(seed, n_train, n_test):
 18    rng = np.random.default_rng(seed)
 19    rho = 0.9
 20    def ar(n):
 21        x = np.empty((n, 10), dtype=np.float32)
 22        x[0] = rng.normal(size=10)
 23        q = math.sqrt(1.0 - rho * rho)
 24        for i in range(1, n):
 25            x[i] = rho * x[i-1] + q * rng.normal(size=10)
 26        return x
 27    xtr, xte = ar(n_train), ar(n_test)
 28    def yfun(x):
 29        y = (1.5*x[:, 0] - 1.1*x[:, 1] + .7*x[:, 2]**2
 30             + .4*np.sin(x[:, 3]) + .25*x[:, 4]*x[:, 5]
 31             + .15*rng.normal(size=len(x)))
 32        return y.astype(np.float32).reshape(-1, 1)
 33    return {'xtr': xtr, 'ytr': yfun(xtr), 'xte': xte, 'yte': yfun(xte),
 34            'task': 'regression', 'metric': 'mse', 'out_dim': 1}
 35
 36
 37def seed_all(seed):
 38    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 39    if torch.cuda.is_available():
 40        torch.cuda.manual_seed_all(seed)
 41
 42
 43def flat_grads(net):
 44    return torch.cat([p.grad.detach().reshape(-1) for p in net.parameters() if p.grad is not None])
 45
 46
 47def per_sample_grad(net, x, y):
 48    vals = []
 49    lossf = nn.MSELoss()
 50    for j in range(len(x)):
 51        net.zero_grad(set_to_none=True)
 52        lossf(net(x[j:j+1]), y[j:j+1]).backward()
 53        vals.append(flat_grads(net))
 54    return torch.stack(vals)
 55
 56
 57def clip(v, bound):
 58    n = torch.linalg.vector_norm(v)
 59    return v if float(n) <= bound else v * (bound / (n + 1e-12))
 60
 61
 62def train_one(ds, seed, lr, method, epochs=8, block=32, levels=2, bound=10.0):
 63    seed_all(seed)
 64    net = make_model('mlp_tiny', (10,), 1)
 65    device = 'cuda' if torch.cuda.is_available() else 'cpu'
 66    try:
 67        net = net.to(device)
 68        x = torch.as_tensor(ds['xtr'], dtype=torch.float32, device=device)
 69        y = torch.as_tensor(ds['ytr'], dtype=torch.float32, device=device)
 70        xt = torch.as_tensor(ds['xte'], dtype=torch.float32, device=device)
 71        yt = torch.as_tensor(ds['yte'], dtype=torch.float32, device=device)
 72        opt = torch.optim.Adam(net.parameters(), lr=lr)
 73        rng = np.random.default_rng(seed + 991)
 74        p = np.array([.5, .3, .2], dtype=float)
 75        costs = []
 76        corr_norms, grad_seq, clipped = [], [], 0
 77        cursor = 0
 78        base = block * (2 ** levels)
 79        for _ in range(epochs):
 80            while cursor + base <= len(x):
 81                z = x[cursor:cursor+base]; q = y[cursor:cursor+base]; cursor += base
 82                if method == 'baseline':
 83                    net.zero_grad(set_to_none=True)
 84                    loss = ((net(z) - q) ** 2).mean()
 85                    loss.backward(); gh = flat_grads(net); costs.append(base)
 86                else:
 87                    # Shared-prefix coupled multilevel estimator: g0 plus one inverse-probability correction.
 88                    l = int(rng.choice(3, p=p))
 89                    b0 = block
 90                    net.zero_grad(set_to_none=True)
 91                    g0 = per_sample_grad(net, z[:b0], q[:b0]).mean(0)
 92                    if l == 0:
 93                        gh = g0; costs.append(b0)
 94                    else:
 95                        bf = b0 * (2 ** l); bp = bf // 2
 96                        gf = per_sample_grad(net, z[:bf], q[:bf]).mean(0)
 97                        gp = per_sample_grad(net, z[:bp], q[:bp]).mean(0)
 98                        delta = gf - gp
 99                        corr_norms.append(float(torch.linalg.vector_norm(delta)))
100                        delta = clip(delta, bound)
101                        if float(torch.linalg.vector_norm(delta)) >= bound - 1e-6: clipped += 1
102                        gh = g0 + delta / float(p[l]); costs.append(bf)
103                    gh = clip(gh, bound)
104                    if float(torch.linalg.vector_norm(gh)) >= bound - 1e-6: clipped += 1
105                net.zero_grad(set_to_none=True)
106                off = 0
107                for par in net.parameters():
108                    n = par.numel(); par.grad = gh[off:off+n].reshape_as(par).clone(); off += n
109                opt.step(); grad_seq.append(gh.detach().cpu().numpy())
110            cursor = 0
111        net.eval()
112        with torch.no_grad(): metric = float(((net(xt)-yt)**2).mean())
113        a = np.asarray(grad_seq)
114        ac = float(np.corrcoef(a[:-1,0], a[1:,0])[0,1]) if len(a) > 3 else float('nan')
115        return metric, net, {'grad_lag1': ac, 'correction_norm_mean': float(np.mean(corr_norms)) if corr_norms else 0.0,
116                             'clip_fraction': clipped/max(1, len(grad_seq)), 'updates': len(grad_seq),
117                             'mean_grad_norm': float(np.linalg.norm(a, axis=1).mean())}
118    except RuntimeError:
119        if device == 'cuda':
120            torch.cuda.empty_cache()
121            return train_one_cpu(ds, seed, lr, method, epochs, block, levels, bound)
122        raise
123
124
125def train_one_cpu(ds, seed, lr, method, epochs=8, block=32, levels=2, bound=10.0):
126    old = torch.cuda.is_available
127    torch.cuda.is_available = lambda: False
128    try: return train_one(ds, seed, lr, method, epochs, block, levels, bound)
129    finally: torch.cuda.is_available = old
130
131
132def run():
133    lrs = [0.001, 0.003, 0.006]
134    seeds = tuple(range(8))
135    def base_fn(cfg):
136        lr = cfg['lr']
137        return lambda seed: train_one(get_dataset(seed, 768, 256), seed, lr, 'baseline')[0]
138    base = sweep_baseline(base_fn, [{'lr': v} for v in lrs], seeds=seeds[:4])
139    def idea_fn(cfg):
140        lr = cfg['lr']
141        return lambda seed: train_one(get_dataset(seed, 768, 256), seed, lr, 'coupled')[0]
142    # Evaluate the baseline-best setting and two nearby settings; all are in the baseline union.
143    idea_trials = []
144    for cfg in [{'lr': v} for v in lrs]:
145        res = evaluate(idea_fn(cfg), seeds=seeds)
146        idea_trials.append({'cfg': cfg, 'result': res})
147    idea = min(idea_trials, key=lambda z: z['result']['mean'])['result']
148    idea['sweep'] = [{'cfg': z['cfg'], 'mean': z['result']['mean'], 'std': z['result']['std']}
149                     for z in idea_trials]
150    sig = []
151    chosen_lr = min(idea_trials, key=lambda z: z['result']['mean'])['cfg']['lr']
152    for seed in seeds:
153        ds = get_dataset(seed, 768, 256)
154        m, _, s = train_one(ds, seed, chosen_lr, 'coupled')
155        sig.append({'seed': seed, 'metric': m, **s})
156    signature = {
157        'prediction': 'shared fine-minus-coarse corrections have smaller norm than raw fine gradients under correlated streams',
158        'observed_correction_norm_mean': float(np.mean([r['correction_norm_mean'] for r in sig])),
159        'observed_mean_gradient_norm': float(np.mean([r['mean_grad_norm'] for r in sig])),
160        'observed_gradient_lag1_mean': float(np.mean([r['grad_lag1'] for r in sig])),
161        'observed_clip_fraction_mean': float(np.mean([r['clip_fraction'] for r in sig])),
162        'confirmed': bool(np.mean([r['correction_norm_mean'] for r in sig]) < np.mean([r['mean_grad_norm'] for r in sig]))
163    }
164    rep = make_report('markov_stream_regression', 'mlp_tiny', base, idea, signature)
165    rep['custom_track'] = {'name': META['name'], 'file': 'markov_stream_bench.py', 'domain': META['domain']}
166    os.makedirs('artifacts', exist_ok=True)
167    with open('artifacts/bench_report.json', 'w') as f: json.dump(rep, f, indent=2)
168    print(json.dumps(rep, indent=2))
169
170if __name__ == '__main__': run()