Discounted Saddle-Gap Controller / experiment.py
Failed on benchmark
1import numpy as np
2
3
4def discounted(gaps, rho):
5 r = 0.0
6 out = []
7 for g in gaps:
8 r = rho * r + (1.0 - rho) * g
9 out.append(r)
10 return np.asarray(out)
11
12
13def exact_gap(x, y, a, box=2.0):
14 # f(x,y)=a*x*y, with x,y in [-box,box].
15 return 2.0 * box * abs(a) * (abs(x) + abs(y))
16
17
18def probe_gap(x, y, a, k=3, eta=0.08, box=2.0):
19 yp, xp = float(y), float(x)
20 for _ in range(k):
21 yp = np.clip(yp + eta * a * x, -box, box)
22 for _ in range(k):
23 xp = np.clip(xp - eta * a * y, -box, box)
24 return a * x * yp - a * xp * y
25
26
27def verify_math():
28 rng = np.random.default_rng(4)
29 rho = 0.9
30 gaps = rng.uniform(0, 3, 80)
31 direct = discounted(gaps, rho)
32 rec = np.zeros_like(gaps)
33 for t, g in enumerate(gaps):
34 rec[t] = rho * (rec[t - 1] if t else 0.0) + (1.0 - rho) * g
35 recursion_err = float(np.max(np.abs(direct - rec)))
36 impulse = discounted(np.r_[1.0, np.zeros(30)], rho)
37 expected = impulse[0] * rho ** np.arange(1, 31)
38 decay_err = float(np.max(np.abs(impulse[1:] - expected)))
39 return recursion_err, decay_err, 1.0 / (1.0 - rho)
40
41
42def run(controller, seed=7, steps=1800, change=650):
43 rng = np.random.default_rng(seed)
44 x, y = 0.9, -0.7
45 ex = ey = 0.055
46 rho, tau = 0.9, 0.025
47 R = 0.0
48 previous = None
49 down_count = 0
50 up_count = 0
51 gaps, norms, rates, switches = [], [], [], []
52 for t in range(steps):
53 a = 1.0 if t < change else -1.0
54 # Mild deterministic measurement noise models minibatch gap noise.
55 noisy_a = a * (1.0 + 0.015 * rng.normal())
56 x_old, y_old = x, y
57 # Simultaneous descent/ascent on the current bilinear payoff.
58 x = np.clip(x - ex * noisy_a * y_old, -2.0, 2.0)
59 y = np.clip(y + ey * noisy_a * x_old, -2.0, 2.0)
60 g = max(0.0, probe_gap(x, y, noisy_a))
61 R = rho * R + (1.0 - rho) * g
62 if controller and previous is not None:
63 if R > previous * (1.0 + tau):
64 ex *= 0.5; ey *= 0.5
65 down_count += 1; up_count = 0
66 elif R < previous * (1.0 - tau):
67 up_count += 1
68 if up_count >= 4:
69 ex = min(0.09, ex * 1.05); ey = min(0.09, ey * 1.05)
70 up_count = 0
71 else:
72 up_count = 0
73 previous = R
74 gaps.append(exact_gap(x, y, a))
75 norms.append(abs(x) + abs(y))
76 rates.append(ex)
77 switches.append(down_count)
78 arr = np.asarray(gaps)
79 return {
80 'mean_gap_all': float(arr.mean()),
81 'mean_gap_after_change': float(arr[change:].mean()),
82 'mean_gap_last_300': float(arr[-300:].mean()),
83 'max_gap_after_change': float(arr[change:].max()),
84 'final_norm': float(norms[-1]),
85 'step_min': float(min(rates)),
86 'step_final': float(rates[-1]),
87 'down_events': int(down_count),
88 }
89
90
91if __name__ == '__main__':
92 rec_err, decay_err, horizon = verify_math()
93 print('MATH recursion_max_abs_error', rec_err)
94 print('MATH impulse_decay_max_abs_error', decay_err)
95 print('MATH effective_horizon', horizon)
96 for controller in (False, True):
97 vals = [run(controller, seed=s) for s in (7, 8, 9)]
98 print('CONTROLLER', controller)
99 for key in vals[0]:
100 print(key, np.mean([v[key] for v in vals]), '+/-', np.std([v[key] for v in vals]))