Support-Scenario Attention Pruning / support_pruning.py

Mechanism failed

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