Polar Slack Attention / run_experiment.py

Mechanism failed

Raw ⬇ ZIP
 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()