Oracle symmetry-level selection / experiment.py
Mechanism failed
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()