File size: 3,685 Bytes
0e5c1d6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
"""Claim 5: Lemma 2.1 (argmax generically a singleton) and Lemma 2.2 (the convex
hull of tokens shrinks under step sizes gamma in (0,1)), underpinning the
Frank-Wolfe reformulation of hardmax attention.

Previously unaddressed: neither lemma was tested.

Hardmax attention as a Frank-Wolfe step: each token moves toward the token that
maximises its inner product,
    j*(i) = argmax_j <x_i, x_j>,      x_i <- (1-gamma) x_i + gamma x_{j*(i)}
which is exactly a Frank-Wolfe step toward a VERTEX of the current convex hull.
"""
import json, numpy as np
from scipy.spatial import ConvexHull

RES = {}


def step(X, gamma):
    S = X @ X.T
    order = np.argsort(-S, axis=1)
    j = order[:, 0]
    top2 = S[np.arange(len(X)), order[:, 0]]-S[np.arange(len(X)), order[:, 1]]
    return (1-gamma)*X+gamma*X[j], top2


def run():
    # ---- Lemma 2.1: argmax is generically a singleton ----
    rows = []
    for d in (2, 3, 5, 10):
        for n in (20, 50):
            gaps, ties = [], 0; tot = 0
            for s in range(40):
                rng = np.random.default_rng(s)
                X = rng.normal(size=(n, d))
                _, g = step(X, 0.3)
                gaps.append(float(np.min(g))); ties += int(np.sum(g <= 1e-12)); tot += n
            rows.append({"d": d, "n": n, "trials": 40, "tokens_checked": tot,
                         "exact_ties": ties, "min_argmax_gap": min(gaps)})
            print("  L2.1 d=%-3d n=%-3d  exact ties %d/%d   smallest argmax gap = %.3e"
                  % (d, n, ties, tot, min(gaps)), flush=True)
    RES["lemma21_singleton"] = {"rows": rows,
        "total_tokens": sum(r["tokens_checked"] for r in rows),
        "total_exact_ties": sum(r["exact_ties"] for r in rows),
        "smallest_gap_overall": min(r["min_argmax_gap"] for r in rows)}

    # ---- Lemma 2.2: convex hull shrinks for gamma in (0,1) ----
    rows = []
    for gamma in (0.1, 0.3, 0.5, 0.9):
        for d in (2, 3):
            vols, diams, mono_v, mono_d, runs = [], [], 0, 0, 0
            for s in range(12):
                rng = np.random.default_rng(100+s)
                X = rng.normal(size=(40, d))
                V, D = [], []
                for t in range(25):
                    try: V.append(ConvexHull(X).volume)
                    except Exception: V.append(0.0)
                    D.append(float(np.max(np.linalg.norm(X[:, None]-X[None, :], axis=-1))))
                    X, _ = step(X, gamma)
                V = np.array(V); D = np.array(D)
                mono_v += int(np.all(np.diff(V) <= 1e-9)); mono_d += int(np.all(np.diff(D) <= 1e-9))
                vols.append(float(V[-1]/max(V[0], 1e-12))); diams.append(float(D[-1]/max(D[0], 1e-12)))
                runs += 1
            rows.append({"gamma": gamma, "d": d, "runs": runs,
                         "volume_monotone_runs": mono_v, "diameter_monotone_runs": mono_d,
                         "final_over_initial_volume": round(float(np.mean(vols)), 6),
                         "final_over_initial_diameter": round(float(np.mean(diams)), 6)})
            print("  L2.2 gamma=%.1f d=%d  volume non-increasing %d/%d, diameter %d/%d | vol ratio %.4f, diam ratio %.4f"
                  % (gamma, d, mono_v, runs, mono_d, runs,
                     rows[-1]["final_over_initial_volume"], rows[-1]["final_over_initial_diameter"]), flush=True)
    RES["lemma22_hull_shrinks"] = {"rows": rows,
        "all_volume_monotone": all(r["volume_monotone_runs"] == r["runs"] for r in rows),
        "all_diameter_monotone": all(r["diameter_monotone_runs"] == r["runs"] for r in rows)}
    json.dump(RES, open("fw_results.json", "w"), indent=1)


if __name__ == "__main__":
    run(); print("DONE")