Oracle symmetry-level selection / experiment.py

Mechanism failed

Raw ⬇ ZIP
  1import json
  2import math
  3from pathlib import Path
  4import numpy as np
  5
  6
  7def shift_batch(x, k):
  8    return np.roll(x, k, axis=1)
  9
 10
 11def orbit_average(x, q):
 12    return sum(shift_batch(x, k) for k in range(q)) / float(q)
 13
 14
 15def ridge_fit(x, y, reg=1e-2):
 16    xb = np.c_[np.ones(len(x)), x]
 17    a = xb.T @ xb + reg * np.eye(xb.shape[1])
 18    a[0, 0] -= reg
 19    return np.linalg.solve(a, xb.T @ y)
 20
 21
 22def predict(w, x):
 23    return np.c_[np.ones(len(x)), x] @ w
 24
 25
 26def mse(w, x, y):
 27    return float(np.mean((predict(w, x) - y) ** 2))
 28
 29
 30def make_data(n, length, asymmetry, noise, seed):
 31    rng = np.random.default_rng(seed)
 32    # Random periodic curves with a low-frequency invariant component.
 33    t = np.arange(length)[None, :]
 34    phase = rng.uniform(0, 2 * np.pi, size=(n, 1))
 35    amp = rng.normal(size=(n, 1))
 36    x = amp * np.cos(2 * np.pi * t / length + phase)
 37    x += 0.35 * rng.normal(size=(n, length))
 38    # Label is mostly rotation invariant, with a controllable absolute-phase term.
 39    invariant = amp[:, 0]
 40    phase_sensitive = x[:, 0]
 41    y = invariant + asymmetry * phase_sensitive + noise * rng.normal(size=n)
 42    return x, y
 43
 44
 45def estimate_A(w, x, y, q):
 46    # Displayed proxy: mean transformed-input loss minus original loss.
 47    original = (predict(w, x) - y) ** 2
 48    transformed = np.stack([(predict(w, shift_batch(x, k)) - y) ** 2 for k in range(q)], 1)
 49    return float(np.mean(transformed) - np.mean(original))
 50
 51
 52def consistency(w, x, q):
 53    p0 = predict(w, x)
 54    ps = np.stack([predict(w, shift_batch(x, k)) for k in range(q)], 1)
 55    return float(np.mean((ps - p0[:, None]) ** 2))
 56
 57
 58def run_case(asymmetry, n_train=160, n_val=160, n_test=160, seed=0):
 59    L = 8
 60    xt, yt = make_data(n_train, L, asymmetry, .20, seed)
 61    xv, yv = make_data(n_val, L, asymmetry, .20, seed + 1)
 62    xe, ye = make_data(n_test, L, asymmetry, .20, seed + 2)
 63    qs = [1, 2, 4, 8]
 64    # Unrestricted pilot estimates the invariance-induced loss discrepancy.
 65    w0 = ridge_fit(xt, yt)
 66    A = {q: max(0.0, estimate_A(w0, xv, yv, q)) for q in qs}
 67    D = {q: consistency(w0, xv, q) for q in qs}
 68    beta = 1.0
 69    n, m, qsharp = n_train, L, 8
 70    scores = {q: (n * m * m * q) ** (-beta / (beta + 1)) + A[q] for q in qs if q <= qsharp}
 71    selected = min(scores, key=scores.get)
 72    models = {}
 73    test_losses = {}
 74    for q in qs:
 75        zt, zv = orbit_average(xt, q), orbit_average(xv, q)
 76        ze = orbit_average(xe, q)
 77        w = ridge_fit(zt, yt)
 78        models[q] = w
 79        test_losses[q] = mse(w, ze, ye)
 80    # Also report a validation-loss selector as a useful practical reference.
 81    val_losses = {q: mse(models[q], orbit_average(xv, q), yv) for q in qs}
 82    return {
 83        'asymmetry': asymmetry, 'A_hat': A, 'D_q': D, 'oracle_scores': scores,
 84        'selected_q': selected, 'test_mse': test_losses, 'val_mse': val_losses,
 85        'q1_test': test_losses[1], 'qmax_test': test_losses[8],
 86        'selected_test': test_losses[selected]
 87    }
 88
 89
 90def math_sanity():
 91    qs = np.array([1, 2, 4, 8], dtype=float)
 92    variance = (160 * 8 * 8 * qs) ** (-1 / 2)
 93    bias = np.array([0.0, .003, .004, .07])
 94    score = variance + bias
 95    assert np.all(np.diff(variance) < 0), variance
 96    assert np.all(np.diff(bias) >= 0), bias
 97    assert int(qs[np.argmin(score)]) == 4
 98    # Directly verify averaging is invariant to every shift in its orbit.
 99    rng = np.random.default_rng(3)
100    x = rng.normal(size=(5, 8))
101    z = orbit_average(x, 4)
102    assert np.max(np.abs(z - shift_batch(z, 1))) > 1e-6  # prefix orbit is not full-group invariant
103    zfull = orbit_average(x, 8)
104    assert np.max(np.abs(zfull - shift_batch(zfull, 1))) < 1e-12
105    return {'variance_proxy': variance.tolist(), 'bias': bias.tolist(), 'score': score.tolist(), 'argmin_q': 4}
106
107
108def main():
109    sanity = math_sanity()
110    cases = [run_case(a, seed=20 + i * 10) for i, a in enumerate([0.0, 0.15, 0.5, 1.0])]
111    out = {'math_sanity': sanity, 'cases': cases}
112    Path('results.json').write_text(json.dumps(out, indent=2))
113    print(json.dumps(out, indent=2))
114
115
116if __name__ == '__main__':
117    main()