Support-Scenario Attention Pruning / support_pruning.py
Mechanism failed
1import json
2import math
3from collections import deque
4import numpy as np
5
6SEED = 376
7
8
9def softmax(x):
10 z = x - np.max(x, axis=-1, keepdims=True)
11 e = np.exp(z)
12 return e / e.sum(axis=-1, keepdims=True)
13
14
15def kappa(p, eps=0.0):
16 return (p > eps).astype(np.int8)
17
18
19def verify_support_order(rng, trials=5000):
20 # Finite Boolean-support version of q <= p iff kappa(q) <= kappa(p).
21 for _ in range(trials):
22 p = rng.random(8)
23 q = rng.random(8)
24 # Make arbitrary nonnegative distributions, including zeros.
25 p[rng.random(8) < .35] = 0
26 q[rng.random(8) < .35] = 0
27 p /= p.sum() if p.sum() else 1
28 q /= q.sum() if q.sum() else 1
29 lhs = np.all((q > 0) <= (p > 0))
30 rhs = np.all(kappa(q) <= kappa(p))
31 if lhs != rhs:
32 return False, _
33 return True, trials
34
35
36def windows(n, width=5, stride=3):
37 out = []
38 start = 0
39 while start < n:
40 out.append(np.arange(start, min(n, start + width)))
41 if start + width >= n:
42 break
43 start += stride
44 return out
45
46
47def independent_mask(prob, eps):
48 return prob > eps
49
50
51def continuation_violations(mask, wins):
52 # A local-surjectivity analogue: in every shared query of neighboring
53 # windows, both retained fibers must share at least one retained key.
54 bad = 0
55 pairs = 0
56 for a, b in zip(wins[:-1], wins[1:]):
57 shared_q = np.intersect1d(a, b)
58 shared_k = np.intersect1d(a, b)
59 for q in shared_q:
60 pairs += 1
61 if not np.any(mask[q, shared_k]):
62 bad += 1
63 return bad, pairs
64
65
66def row_empty(mask):
67 return int(np.sum(mask.sum(axis=1) == 0))
68
69
70def repair_continuations(prob, mask, wins):
71 """Add one maximum-weight common-key witness per violated overlap fiber."""
72 mask = mask.copy()
73 for a, b in zip(wins[:-1], wins[1:]):
74 shared_q = np.intersect1d(a, b)
75 shared_k = np.intersect1d(a, b)
76 for q in shared_q:
77 if not np.any(mask[q, shared_k]):
78 k = shared_k[np.argmax(prob[q, shared_k])]
79 mask[q, k] = True
80 return mask
81
82
83def structured_clean(prob, initial, wins):
84 mask = initial.copy()
85 n = mask.shape[0]
86 # Remove low-weight edges only when all fibers and overlap continuations
87 # remain nonempty. This is the finite mask certificate used by the MVP.
88 candidates = [(float(prob[q, k]), q, k) for q in range(n) for k in range(n)
89 if mask[q, k]]
90 candidates.sort()
91 for _, q, k in candidates:
92 if mask[q].sum() <= 1:
93 continue
94 mask[q, k] = False
95 bad, _ = continuation_violations(mask, wins)
96 if bad or row_empty(mask):
97 mask[q, k] = True
98 return mask
99
100
101def window_graph(mask, wins):
102 g = {i: set() for i in range(len(wins))}
103 for i in range(len(wins)):
104 for j in range(i + 1, len(wins)):
105 shared_q = np.intersect1d(wins[i], wins[j])
106 shared_k = np.intersect1d(wins[i], wins[j])
107 # A bijective restriction in this row-wise toy means each shared
108 # query has exactly one common surviving key.
109 bij = bool(len(shared_q) and all(np.sum(mask[q, shared_k]) == 1
110 for q in shared_q))
111 if bij:
112 g[i].add(j); g[j].add(i)
113 return g
114
115
116def components(g):
117 seen, cs = set(), []
118 for s in g:
119 if s in seen: continue
120 c, todo = [], [s]; seen.add(s)
121 while todo:
122 u = todo.pop(); c.append(u)
123 for v in g[u]:
124 if v not in seen: seen.add(v); todo.append(v)
125 cs.append(c)
126 return cs
127
128
129def attention_error(logits, values, mask, dense):
130 masked = np.where(mask, logits, -1e9)
131 out = softmax(masked) @ values
132 return float(np.mean((out - dense) ** 2))
133
134
135def one_trial(seed, n=12, d=8):
136 rng = np.random.default_rng(seed)
137 wins = windows(n, width=6, stride=3)
138 q = rng.normal(size=(n, d)); k = rng.normal(size=(n, d)); v = rng.normal(size=(n, d))
139 raw = q @ k.T / math.sqrt(d)
140 local = np.zeros((n, n), dtype=bool)
141 for w in wins: local[np.ix_(w, w)] = True
142 logits = np.where(local, raw, -1e9)
143 prob = softmax(logits)
144 eps = 0.08
145 base = prob > eps
146 idea = structured_clean(prob, repair_continuations(prob, base, wins), wins)
147 dense = prob @ v
148 bad0, pairs = continuation_violations(base, wins)
149 bad1, _ = continuation_violations(idea, wins)
150 return {
151 "dense_edges": int(local.sum()), "dense_error": 0.0,
152 "independent_edges": int(base.sum()), "independent_error": attention_error(logits, v, base, dense),
153 "independent_violations": bad0,
154 "structured_edges": int(idea.sum()), "structured_error": attention_error(logits, v, idea, dense),
155 "structured_violations": bad1, "checks": pairs,
156 "independent_components": len(components(window_graph(base, wins))),
157 "structured_components": len(components(window_graph(idea, wins)))
158 }
159
160
161def run():
162 rng = np.random.default_rng(SEED)
163 support_ok, support_trials = verify_support_order(rng)
164 trials = [one_trial(SEED + i) for i in range(10)]
165 keys = trials[0].keys()
166 summary = {}
167 for key in keys:
168 vals = [x[key] for x in trials]
169 summary[key] = {"mean": float(np.mean(vals)), "std": float(np.std(vals)), "values": vals}
170 result = {
171 "seed": SEED, "epsilon": 0.08, "math_support_order": {"passed": support_ok, "trials": support_trials},
172 "trials": 10, "summary": summary
173 }
174 with open("results.json", "w") as f: json.dump(result, f, indent=2)
175 print(json.dumps(result, indent=2))
176
177
178if __name__ == "__main__":
179 run()