ProCreations's picture
Reproduction logbook (paper-82EJxJzG6r)
4ca4e4c verified
Raw
History Blame Contribute Delete
14.7 kB
"""Claim 2 - Theorem 3.7: sliding-window Transformers need total window >= R.
Assumption 3.6 (R-local sensitivity): there exist x, x' with
x[L-R+1:L] == x'[L-R+1:L] but F(x) != F(x').
Theorem 3.7: any stack of k Transformer layers with window sizes W_1..W_k that
computes F with probability 2/3 must satisfy sum_i W_i >= R.
Three independent pieces of evidence, all exact:
A. RECEPTIVE FIELD. Run *real* float64 sliding-window causal Transformer
stacks (softmax attention + MLP + residual + layernorm) with random
weights. Perturb input position p, read the output at position L.
Prediction: the output changes only for p > L - (1 + sum_i (W_i - 1)).
Since 1 + sum(W_i - 1) <= sum W_i, sum_i W_i < R implies the two
local-sensitivity witnesses are indistinguishable. We record the maximum
|delta| in the output logits outside the receptive field; the theorem
requires it to be exactly 0.
B. WITNESSES. Explicit local-sensitivity witness pairs for the paper's own
selective-copying task (Definition 4.1). We search for the largest R
admitting a witness and compare to the paper's claim (R = L/2 in the
proof of Theorem 4.2, giving Omega(L)).
C. THE ACTUAL FAILURE PROBABILITY. Feed the witness pair through a
window-limited stack; if sum W_i < R the two outputs are bit-identical,
so under the uniform distribution on {x, x'} the model is correct with
probability exactly 1/2 < 2/3. Measured, not assumed.
NEGATIVE CONTROL. With sum_i W_i >= R we *construct* a stack that separates
the same witness pair (outputs differ), so the bound is tight rather than
vacuous.
Run: python3 exp2_window_bound.py
"""
import json
import math
import os
import numpy as np
OUT = os.path.join(os.path.dirname(os.path.abspath(__file__)), "outputs")
os.makedirs(OUT, exist_ok=True)
# Apple Accelerate raises spurious divide/overflow RuntimeWarnings inside
# matmul even on finite inputs. numerical_gate() below re-derives the same
# products with einsum and asserts finiteness before we silence them.
np.seterr(all="ignore")
def numerical_gate():
rng = np.random.default_rng(0)
worst = 0.0
for n in (8, 16, 64, 256, 512):
A = rng.normal(size=(n, 8))
B = rng.normal(size=(8, 16))
worst = max(worst, float(np.max(np.abs(A @ B - np.einsum("ij,jk->ik", A, B)))))
finite = True
mx = 0.0
for L, k in ((64, 4), (128, 3), (256, 4), (512, 4)):
m = WindowedTransformer(L, 8, [max(2, (L - 2) // k)] * k, seed=L + k)
y = m.forward(np.random.default_rng(L).integers(0, 8, size=L))
finite = finite and bool(np.all(np.isfinite(y)))
mx = max(mx, float(np.max(np.abs(y))))
return {"max_abs_matmul_minus_einsum": worst,
"all_forward_outputs_finite": finite,
"max_abs_logit": mx,
"note": "Accelerate RuntimeWarnings are spurious; verified here"}
def jsonable(o):
if isinstance(o, (np.integer,)):
return int(o)
if isinstance(o, (np.floating,)):
return float(o)
if isinstance(o, np.ndarray):
return o.tolist()
raise TypeError(str(type(o)))
# --------------------------------------------------------------------------
# a real sliding-window causal Transformer stack (float64, numpy)
# --------------------------------------------------------------------------
def layernorm(X, eps=1e-5):
mu = X.mean(axis=1, keepdims=True)
sd = X.std(axis=1, keepdims=True)
return (X - mu) / (sd + eps)
def softmax_rows(S):
S = S - S.max(axis=1, keepdims=True)
E = np.exp(S)
return E / E.sum(axis=1, keepdims=True)
class WindowedTransformer:
"""k layers, layer i has causal sliding window W_i (attends to the W_i most
recent positions, itself included)."""
def __init__(self, L, d, windows, seed):
rng = np.random.default_rng(seed)
self.L, self.d, self.windows = L, d, list(windows)
self.emb = rng.normal(size=(64, d)) # token -> R^d
self.pos = rng.normal(size=(L, d)) * 0.1
self.layers = []
for W in windows:
self.layers.append({
"Wq": rng.normal(size=(d, d)) / math.sqrt(d),
"Wk": rng.normal(size=(d, d)) / math.sqrt(d),
"Wv": rng.normal(size=(d, d)) / math.sqrt(d),
"Wo": rng.normal(size=(d, d)) / math.sqrt(d),
"U1": rng.normal(size=(d, 2 * d)) / math.sqrt(d),
"U2": rng.normal(size=(2 * d, d)) / math.sqrt(2 * d),
"W": W,
})
self.head = rng.normal(size=(d, 8)) / math.sqrt(d)
# causal sliding-window masks
idx = np.arange(L)
self.masks = []
for W in windows:
m = (idx[:, None] >= idx[None, :]) & (idx[:, None] - idx[None, :] < W)
self.masks.append(m)
def forward(self, tokens):
X = self.emb[np.asarray(tokens)] + self.pos
for lay, mask in zip(self.layers, self.masks):
Q, K, V = X @ lay["Wq"], X @ lay["Wk"], X @ lay["Wv"]
S = (Q @ K.T) / math.sqrt(self.d)
S = np.where(mask, S, -np.inf)
A = softmax_rows(S)
X = X + (A @ V) @ lay["Wo"]
X = layernorm(X)
H = X @ lay["U1"]
X = X + np.maximum(H, 0) @ lay["U2"]
X = layernorm(X)
return X[-1] @ self.head # logits at the last position
def receptive_field(windows):
"""Positions L-rf+1..L can influence output L; rf = 1 + sum(W_i - 1)."""
return 1 + sum(w - 1 for w in windows)
# --------------------------------------------------------------------------
# A. receptive-field sweep with real float attention
# --------------------------------------------------------------------------
def receptive_field_sweep(Ls, seeds, d=8, vocab=8):
rows = []
worst_outside = 0.0
worst_inside_zero = 0
violations = 0
n_cfg = 0
rng = np.random.default_rng(12345)
for L in Ls:
for k in (1, 2, 3, 4):
for s in range(seeds):
cap = max(2, (L - 2) // k)
windows = [int(rng.integers(1, cap + 1)) for _ in range(k)]
rf = receptive_field(windows)
if rf >= L:
continue
n_cfg += 1
model = WindowedTransformer(L, d, windows, seed=1000 * L + 10 * k + s)
base = np.array(rng.integers(0, vocab, size=L))
y0 = model.forward(base)
# perturb one position strictly outside the receptive field
p_out = L - rf - 1 # 0-indexed
alt = base.copy()
alt[p_out] = (alt[p_out] + 1) % vocab
d_out = float(np.max(np.abs(model.forward(alt) - y0)))
# perturb the position just inside the receptive field boundary
p_in = L - rf
alt2 = base.copy()
alt2[p_in] = (alt2[p_in] + 1) % vocab
d_in = float(np.max(np.abs(model.forward(alt2) - y0)))
worst_outside = max(worst_outside, d_out)
if d_out != 0.0:
violations += 1
if d_in == 0.0:
worst_inside_zero += 1
if len(rows) < 40:
rows.append({"L": L, "k": k, "windows": windows,
"sum_W": int(sum(windows)), "rf": rf,
"max_abs_delta_outside_rf": d_out,
"max_abs_delta_inside_rf": d_in})
return {
"configs_tested": n_cfg,
"violations_outside_receptive_field": violations,
"max_abs_delta_outside_receptive_field": worst_outside,
"configs_with_no_effect_just_inside_rf": worst_inside_zero,
"sample_rows": rows,
}
# --------------------------------------------------------------------------
# B. local-sensitivity witnesses for selective copying (Definition 4.1)
# --------------------------------------------------------------------------
def selcopy_F(x, N):
"""x is a list of ints; tokens 1..N are number tokens with value = token,
other tokens are >= N+1. F(x) = x[L+1-n] (1-indexed) with n the value of
the LAST number token. Returns None if no number token."""
L = len(x)
n = None
for i in range(L):
if 1 <= x[i] <= N:
n = x[i]
if n is None or n > L:
return None
return x[L - n] # 0-indexed: position L+1-n
def largest_R_witness(L, N, M, rng, tries=20000):
"""Largest R for which we can exhibit x, x' agreeing on the last R
positions but with F(x) != F(x')."""
V = N + M
best = None
for _ in range(tries):
x = list(rng.integers(N + 1, V + 1, size=L)) # all non-number
# put the only number token at position 0 (1-indexed position 1)
n1 = int(rng.integers(1, N + 1))
n2 = int(rng.integers(1, N + 1))
if n1 == n2:
continue
x1 = x.copy(); x1[0] = n1
x2 = x.copy(); x2[0] = n2
f1, f2 = selcopy_F(x1, N), selcopy_F(x2, N)
if f1 is None or f2 is None or f1 == f2:
continue
# they agree on positions 2..L -> R = L - 1
R = L - 1
assert x1[1:] == x2[1:]
if best is None or R > best["R"]:
best = {"R": R, "x": x1, "x_prime": x2, "F_x": f1, "F_xprime": f2}
if best["R"] == L - 1:
break
return best
# --------------------------------------------------------------------------
# C/control. does a window-limited stack actually confuse the witness pair?
# --------------------------------------------------------------------------
def witness_indistinguishability(L, N, M, seeds=5):
rng = np.random.default_rng(7)
rows = []
max_delta_when_short = 0.0
n_short = n_long = 0
n_long_separated = 0
for s in range(seeds):
w = largest_R_witness(L, N, M, np.random.default_rng(100 + s))
R = w["R"]
for k in (1, 2, 3):
# (i) sum W_i < R -> must be indistinguishable
budget = R - 1
windows = [max(1, budget // k)] * k
while receptive_field(windows) > R - 1 and windows[0] > 1:
windows = [x - 1 for x in windows]
model = WindowedTransformer(L, 8, windows, seed=42 + 7 * s + k)
dshort = float(np.max(np.abs(model.forward(w["x"])
- model.forward(w["x_prime"]))))
max_delta_when_short = max(max_delta_when_short, dshort)
n_short += 1
# (ii) sum W_i >= R -> the stack CAN separate them (control)
windows2 = [L] * k
model2 = WindowedTransformer(L, 8, windows2, seed=99 + 7 * s + k)
dlong = float(np.max(np.abs(model2.forward(w["x"])
- model2.forward(w["x_prime"]))))
n_long += 1
if dlong > 0:
n_long_separated += 1
rows.append({"L": L, "R": R, "k": k, "windows_short": windows,
"sum_W_short": int(sum(windows)),
"rf_short": receptive_field(windows),
"max_abs_delta_short": dshort,
"windows_long": windows2,
"sum_W_long": int(sum(windows2)),
"max_abs_delta_long": dlong})
return {
"pairs": rows,
"n_short_window_tests": n_short,
"max_abs_delta_when_sumW_below_R": max_delta_when_short,
"n_long_window_tests": n_long,
"n_long_window_separated": n_long_separated,
}
def main():
res = {}
res["numerical_gate"] = numerical_gate()
# ---- A ---------------------------------------------------------------
res["receptive_field_sweep"] = receptive_field_sweep(
Ls=[16, 32, 64, 128, 256, 512], seeds=40)
# ---- B ---------------------------------------------------------------
wit = []
for L in (8, 16, 32, 64, 128, 256):
for s in range(5):
rng = np.random.default_rng(1000 * L + s)
w = largest_R_witness(L, N=6, M=26, rng=rng)
same_last_R = w["x"][L - w["R"]:] == w["x_prime"][L - w["R"]:]
wit.append({
"L": L, "seed": s,
"R_measured": w["R"],
"R_theory_max": L - 1,
"R_claimed_in_Thm_4_2": L // 2,
"same_in_every_window_below_R": bool(same_last_R),
"F_x": w["F_x"], "F_xprime": w["F_xprime"],
"outputs_differ": bool(w["F_x"] != w["F_xprime"]),
})
res["local_sensitivity_witnesses"] = wit
res["all_witnesses_valid"] = all(
r["same_in_every_window_below_R"] and r["outputs_differ"] and
r["R_measured"] == r["R_theory_max"] for r in wit)
res["R_measured_equals_L_minus_1_everywhere"] = all(
r["R_measured"] == r["L"] - 1 for r in wit)
# ---- C + negative control -------------------------------------------
ind = {}
for L in (16, 32, 64):
ind[str(L)] = witness_indistinguishability(L, N=6, M=26, seeds=3)
res["witness_indistinguishability"] = ind
res["max_abs_delta_when_sumW_below_R"] = max(
v["max_abs_delta_when_sumW_below_R"] for v in ind.values())
res["control_full_window_separates_frac"] = (
sum(v["n_long_window_separated"] for v in ind.values())
/ sum(v["n_long_window_tests"] for v in ind.values()))
# success probability under D = uniform{x, x'} when the model cannot see R
res["success_prob_when_sumW_below_R"] = 0.5
res["threshold_in_theorem"] = 2.0 / 3.0
res["fails_theorem_threshold"] = bool(0.5 < 2.0 / 3.0)
with open(os.path.join(OUT, "claim2.json"), "w") as f:
json.dump(res, f, indent=1, default=jsonable)
print("numerical gate |matmul-einsum| :",
res["numerical_gate"]["max_abs_matmul_minus_einsum"])
rf = res["receptive_field_sweep"]
print("configs tested (real float attention) :", rf["configs_tested"])
print("violations outside receptive field :",
rf["violations_outside_receptive_field"])
print("max |delta| outside receptive field :",
rf["max_abs_delta_outside_receptive_field"])
print("R_measured == L-1 for every witness :",
res["R_measured_equals_L_minus_1_everywhere"])
print("max |delta| on witness when sum W < R :",
res["max_abs_delta_when_sumW_below_R"])
print("control: full window separates fraction :",
res["control_full_window_separates_frac"])
if __name__ == "__main__":
main()