Polar Slack Attention / run_experiment.py
Mechanism failed
1import json
2import numpy as np
3from polar_slack import tetrahedron_design, slack_matrix, polar_attention_logits, row_softmax
4
5
6def numerical_rank(A, tol=1e-9):
7 s = np.linalg.svd(A, compute_uv=False)
8 return int(np.sum(s > tol * max(s[0], 1.0))), s
9
10
11def main():
12 np.random.seed(7)
13 X = tetrahedron_design()
14 d, N = X.shape
15 h = np.ones(N)
16 U = -np.eye(d)
17 c = 1.0 / 3.0
18 A = slack_matrix(X, h, c, U)
19 rank, singular = numerical_rank(A)
20 B = X.T @ U.T @ X
21 zero_tol = 1e-10
22
23 # Exact tetrahedral self-polar realization checks.
24 math = {
25 "shape": list(A.shape),
26 "min_A": float(A.min()),
27 "max_abs_A_minus_expected": float(np.max(np.abs(A - (4.0 / 3.0) * np.eye(N)))),
28 "rank_A": rank,
29 "expected_rank_d_plus_1": d + 1,
30 "rank_B": numerical_rank(B)[0],
31 "B_nonzero_singular_values": [float(v) for v in np.linalg.svd(B, compute_uv=False) if v > 1e-9],
32 "row_sums": [float(v) for v in A.sum(1)],
33 "zero_entries_per_row": int(np.sum(np.abs(A[0]) <= zero_tol)),
34 }
35
36 # Scope check: arbitrary supports need not preserve the paper's guarantee.
37 A_invalid = slack_matrix(X, np.array([0.2, 1.0, 1.0, 1.0]), c, U)
38 invalid_check = {"min_A_with_arbitrary_h": float(A_invalid.min()),
39 "has_negative_entry": bool(np.any(A_invalid < -zero_tol))}
40
41 # Controlled attention proxy: each position's geometric incidence target is itself.
42 # Every method receives exactly the same random nuisance QK logits.
43 trials, beta = 2000, 1.0
44 noise = np.random.normal(size=(trials, N, N))
45 qk = beta * noise
46 target = np.broadcast_to(np.arange(N), (trials, N))
47 dense_p = row_softmax(qk)
48 top_idx = np.argmax(qk, axis=-1)
49 top_acc = float(np.mean(top_idx == target))
50 geo_logits = polar_attention_logits(qk, A, alpha=1.0, eps=1e-8)
51 geo_p = row_softmax(geo_logits)
52 geo_idx = np.argmax(geo_logits, axis=-1)
53 geo_acc = float(np.mean(geo_idx == target))
54 rows = np.arange(N)[None, :]
55 dense_target = float(np.mean(dense_p[np.arange(trials)[:, None], rows, target]))
56 geo_target = float(np.mean(geo_p[np.arange(trials)[:, None], rows, target]))
57 experiment = {
58 "trials": trials, "N": N, "beta": beta,
59 "dense_target_probability": dense_target,
60 "polar_target_probability": geo_target,
61 "ordinary_top1_accuracy": top_acc,
62 "polar_argmax_accuracy": geo_acc,
63 "dense_edges_per_row": N,
64 "ordinary_top1_edges_per_row": 1,
65 "polar_positive_edges_per_row": int(np.sum(A > zero_tol)),
66 "polar_effective_edges_per_row": int(np.sum(A[0] > zero_tol)),
67 "polar_cross_entropy": float(-np.mean(np.log(np.maximum(geo_p[np.arange(trials)[:, None], rows, target], 1e-30)))),
68 "dense_cross_entropy": float(-np.mean(np.log(np.maximum(dense_p[np.arange(trials)[:, None], rows, target], 1e-30)))),
69 }
70 print(json.dumps({"math": math, "invalid_relaxation": invalid_check, "experiment": experiment}, indent=2))
71
72
73if __name__ == '__main__':
74 main()