Gauge-Free Inverse OT Attention / experiment.py

Failed on benchmark

Raw ⬇ ZIP
 1import json
 2import numpy as np
 3
 4
 5def center(x):
 6    return x - x.mean(0, keepdims=True) - x.mean(1, keepdims=True) + x.mean()
 7
 8
 9def sinkhorn(cost, eps, s, r, iters):
10    logk = -cost / eps
11    lu = np.zeros(len(s)); lv = np.zeros(len(r))
12    for _ in range(iters):
13        lu = np.log(s) - np.logaddexp.reduce(logk + lv[None, :], axis=1)
14        lv = np.log(r) - np.logaddexp.reduce(logk.T + lu[None, :], axis=1)
15    return np.exp(lu[:, None] + logk + lv[None, :])
16
17
18def inverse_cost(w, eps, delta=0.0):
19    return -eps * center(np.log(w + delta))
20
21
22def math_checks(seed=7):
23    rng = np.random.default_rng(seed)
24    m, n, eps = 9, 7, 0.7
25    c = rng.normal(size=(m, n))
26    s = rng.random(m); s /= s.sum()
27    r = rng.random(n); r /= r.sum()
28    w = sinkhorn(c, eps, s, r, 300)
29    d = np.linalg.norm(center(c))
30    exact = np.linalg.norm(center(c) - inverse_cost(w, eps)) / d
31    marg = max(np.max(abs(w.sum(1)-s)), np.max(abs(w.sum(0)-r)))
32    g = rng.normal(size=m)[:, None] + rng.normal(size=n)[None, :]
33    wg = sinkhorn(c + g, eps, s, r, 300)
34    gauge_plan = np.linalg.norm(w-wg) / np.linalg.norm(w)
35    gauge_inverse = np.linalg.norm(inverse_cost(wg, eps)-inverse_cost(w, eps)) / d
36    convergence = []
37    for it in [3, 5, 10, 30, 100]:
38        wi = sinkhorn(c, eps, s, r, it)
39        convergence.append({
40            'iters': it,
41            'marginal_error': float(max(np.max(abs(wi.sum(1)-s)), np.max(abs(wi.sum(0)-r)))),
42            'recovery_error': float(np.linalg.norm(center(c)-inverse_cost(wi, eps))/d),
43            'min_w': float(wi.min())})
44    floor = []
45    for delta in [0.0, 1e-12, 1e-8, 1e-5, 1e-3]:
46        floor.append({'delta': delta, 'recovery_error': float(np.linalg.norm(center(c)-inverse_cost(w, eps, delta))/d)})
47    return {'exact_recovery_error': float(exact), 'marginal_error': float(marg),
48            'gauge_plan_relative_change': float(gauge_plan),
49            'gauge_inverse_relative_change': float(gauge_inverse),
50            'convergence': convergence, 'floor': floor, 'min_w_300': float(w.min())}
51
52
53def attention_benchmark(seed=11, trials=300):
54    rng = np.random.default_rng(seed)
55    n, d, eps = 12, 16, 0.5
56    methods = ['row_softmax', 'sinkhorn', 'sinkhorn_inverse_penalty']
57    scores = {x: [] for x in methods}; entropy = {x: [] for x in methods}
58    for _ in range(trials):
59        q = rng.normal(size=(n, d)); k = rng.normal(size=(n, d)); v = rng.normal(size=(n, d))
60        target = np.argmax(q @ k.T / np.sqrt(d), axis=1)
61        raw = q @ k.T / np.sqrt(d)
62        row = np.exp(raw-raw.max(1,keepdims=True)); row /= row.sum(1,keepdims=True)
63        w = sinkhorn(-raw, eps, np.ones(n)/n, np.ones(n)/n, 30)
64        # The inverse-cost consistency term is evaluated, not optimized: it is zero
65        # for a plan generated by the same kernel, isolating geometry rather than tuning.
66        inv = inverse_cost(w, eps)
67        structured = center(raw)
68        penalty = np.mean((structured - inv)**2)
69        for name, a in [('row_softmax', row), ('sinkhorn', w), ('sinkhorn_inverse_penalty', w)]:
70            pred = np.argmax(a, axis=1)
71            scores[name].append(np.mean(pred == target))
72            entropy[name].append(float(-np.mean(np.sum(a*np.log(a+1e-30),1))))
73    return {m: {'retrieval_accuracy': float(np.mean(scores[m])),
74                'entropy': float(np.mean(entropy[m]))} for m in methods}
75
76
77if __name__ == '__main__':
78    out = {'math': math_checks(), 'attention': attention_benchmark()}
79    print(json.dumps(out, indent=2))