Tempered-Stable Volatility Clock for Sequence Diffusion / bench_stage2.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, math, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import sys
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
  8
  9SEEDS = tuple(range(8))
 10TRACK = 'sequence'
 11MODEL = 'transformer_tiny'
 12EPOCHS = 5
 13NTRAIN, NTEST = 400, 400
 14# Fixed a priori clock parameters; these produce visible but not extreme clustering.
 15ALPHA, THETA, PHI = 0.65, 0.8, 0.7
 16
 17def positive_stable(alpha, rng, size):
 18    u = rng.uniform(1e-8, math.pi - 1e-8, size)
 19    e = rng.exponential(1.0, size)
 20    return (np.sin(alpha*u) / np.sin(u)**(1/alpha) *
 21            (np.sin((1-alpha)*u) / e)**((1-alpha)/alpha))
 22
 23def ts_sample(alpha, theta, delta, rng, size):
 24    scale = delta ** (1/alpha)
 25    out = np.empty(size)
 26    filled = 0
 27    accept = max(math.exp(-delta * theta**alpha), .05)
 28    while filled < size:
 29        n = max(64, int((size-filled) / accept * 1.15))
 30        x = scale * positive_stable(alpha, rng, n)
 31        keep = rng.random(n) < np.exp(-theta*x)
 32        got = x[keep]
 33        take = min(len(got), size-filled)
 34        if take:
 35            out[filled:filled+take] = got[:take]
 36            filled += take
 37    return out
 38
 39def clock(shape, rng):
 40    delta = (1-PHI) * THETA**(1-ALPHA) / ALPHA
 41    a = np.ones(shape[0])
 42    for _ in range(80):
 43        a = PHI*a + ts_sample(ALPHA, THETA, delta, rng, shape[0])
 44    out = np.empty(shape)
 45    for j in range(shape[1]):
 46        a = PHI*a + ts_sample(ALPHA, THETA, delta, rng, shape[0])
 47        out[:, j] = a
 48    return out.astype(np.float32)
 49
 50def make_ds(seed, idea):
 51    d = get_dataset(TRACK, seed, NTRAIN, NTEST)
 52    if idea:
 53        # Blind exposure: the model receives only the perturbed sequence, not A.
 54        rng = np.random.default_rng(100000 + seed)
 55        A = clock(tuple(d['xtr'].shape), rng)
 56        d = dict(d)
 57        d['xtr'] = d['xtr'] * torch.from_numpy(np.sqrt(A))
 58    return d
 59
 60def run_one(seed, cfg, idea, return_model=False):
 61    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 62    d = make_ds(seed, idea)
 63    net = make_model(MODEL, d['input_shape'], d['out_dim'])
 64    net, metric, hist = train_model(net, d, epochs=EPOCHS, lr=cfg['lr'],
 65                                    batch=128, weight_decay=cfg['weight_decay'], log=lambda *_: None)
 66    if net is None: return float('nan') if not return_model else (None, d)
 67    return (float(metric), net, d) if return_model else float(metric)
 68
 69def factory(idea):
 70    return lambda cfg: lambda seed: run_one(seed, cfg, idea)
 71
 72def prediction_signature(cfg):
 73    # Test behavior of each trained NN under two independent clock perturbations.
 74    vals=[]; observed_A=[]; observed_r=[]
 75    for seed in SEEDS:
 76        got = run_one(seed, cfg, True, True)
 77        if got[0] is None: continue
 78        _, net, d = got
 79        net.eval()
 80        device = next(net.parameters()).device
 81        rng = np.random.default_rng(900000 + seed)
 82        x = d['xte'][:128].to(device)
 83        a = clock(tuple(x.shape), rng)
 84        b = clock(tuple(x.shape), rng)
 85        with torch.no_grad():
 86            p0 = net(x)
 87            p1 = net(x * torch.from_numpy(np.sqrt(a)).to(device))
 88            p2 = net(x * torch.from_numpy(np.sqrt(b)).to(device))
 89        # Prediction response is the trained model's behavior; report its
 90        # perturbation kurtosis and cross-perturbation correlation.
 91        r1 = (p1-p0).detach().cpu().numpy().ravel(); r2 = (p2-p0).detach().cpu().numpy().ravel()
 92        m2 = np.mean(r1*r1); k = np.mean(r1**4)/(m2*m2)-3 if m2 > 1e-12 else 0.
 93        corr = float(np.corrcoef(r1*r1, r2*r2)[0,1]) if np.std(r1*r1)>1e-12 and np.std(r2*r2)>1e-12 else 0.
 94        vals.append({'seed': seed, 'prediction_response_excess_kurtosis': float(k), 'response_sq_corr': corr})
 95        observed_A.append(a); observed_r.append(a.astype(np.float64))
 96    # Clock statistics are measured on perturbations actually fed through the
 97    # trained systems, while the NN response statistics above are independent.
 98    aa=np.concatenate([x.ravel() for x in observed_A]); rr=np.concatenate([x[:,:-1].ravel() for x in observed_r]); ss=np.concatenate([x[:,1:].ravel() for x in observed_r])
 99    v=float(np.var(aa)); K=float(3*v); rho=float(np.corrcoef((rr-1)**2,(ss-1)**2)[0,1])
100    pred_v=(1-ALPHA)/(THETA*(1+PHI)); pred_rho=PHI*pred_v/(2+3*pred_v)
101    # Quantitative confirmation requires the trained response to show the
102    # claimed statistic; this conservative test should not claim confirmation.
103    response_k=float(np.mean([x['prediction_response_excess_kurtosis'] for x in vals])) if vals else float('nan')
104    return {'predicted_clock_excess_kurtosis':3*pred_v, 'observed_clock_excess_kurtosis':K,
105            'predicted_clock_lag1_squared_acf':pred_rho, 'observed_clock_lag1_squared_acf':rho,
106            'trained_model_response': vals, 'response_mean_excess_kurtosis':response_k,
107            'confirmed': False}
108
109def main():
110    grid=[{'lr':lr,'weight_decay':wd} for lr in (0.0015,0.003,0.006) for wd in (0.0,1e-4)]
111    base=sweep_baseline(factory(False), grid, seeds=(0,1,2,3))
112    best=base['best_cfg']
113    # Three idea settings, all lr/wd values included in baseline grid.
114    idea_grid=[best, {'lr':0.0015,'weight_decay':best['weight_decay']}, {'lr':0.006,'weight_decay':best['weight_decay']}]
115    idea_trials=[]
116    for cfg in idea_grid:
117        r=evaluate(factory(True)(cfg), seeds=SEEDS)
118        idea_trials.append({'cfg':cfg,'result':r})
119    chosen=min(idea_trials, key=lambda z:z['result']['mean'])
120    sig=prediction_signature(chosen['cfg'])
121    rep=make_report(TRACK, MODEL, base, chosen['result'], sig)
122    rep['idea_sweep']=idea_trials
123    rep['protocol_notes']={'structural_match':'sequence windows have multi-position correlations; blind clock augmentation is the only intervention', 'epochs':EPOCHS, 'n_train':NTRAIN, 'n_test':NTEST, 'clock':{'alpha':ALPHA,'theta':THETA,'phi':PHI}}
124    Path('bench_report.json').write_text(json.dumps(rep, indent=2, allow_nan=False))
125    print(json.dumps(rep, indent=2))
126if __name__=='__main__': main()