kpshinnik's picture
download
raw
4.25 kB
"""
Claim 4 — numerical verification of DANCE's two theorems (Sec 5).
Theorem 5.4 (bounded distortion from hard truncation):
||t~_v - t_v^full||_2 <= 2 M * delta_Btok(p_v)
where e chunk embeddings bounded ||e||<=M, p_v the untruncated weights,
pi_v = Pi_Btok(p_v), and delta = tail mass beyond top-Btok.
Theorem 5.6 (selection stability under bounded model drift):
if ||omega - omega'||_2 <= Delta_B(omega) / (2 L_s), then the selected top-B
index set (after entmax + top-B truncation) is invariant. L_s is the score
Lipschitz constant.
Both are checked over many random trials; we report the maximum bound-violation
(should be <= 1 for 5.4) and the invariance rate inside vs outside the 5.6 ball.
"""
import json, os, sys
import numpy as np
sys.path.insert(0, os.path.dirname(__file__))
from dance_core import hard_budget_project, tail_mass, softmax, topB_margin, topB_set
rng = np.random.default_rng(0)
def check_thm54(trials=20000):
max_ratio = 0.0
viol = 0
ratios = []
for _ in range(trials):
C = rng.integers(4, 60)
d = rng.integers(4, 64)
B = rng.integers(1, C)
M = float(rng.uniform(0.5, 5.0))
e = rng.normal(size=(C, d))
norms = np.linalg.norm(e, axis=1, keepdims=True)
# scale so max norm <= M (bounded chunk embeddings, Assumption D.1)
e = e / norms.max() * M
p = softmax(rng.normal(size=C) * rng.uniform(0.3, 3.0))
pi = hard_budget_project(p, B)
t_full = (p[:, None] * e).sum(0)
t_trunc = (pi[:, None] * e).sum(0)
lhs = float(np.linalg.norm(t_trunc - t_full))
rhs = 2 * M * tail_mass(p, B)
ratio = lhs / (rhs + 1e-12)
ratios.append(ratio)
max_ratio = max(max_ratio, ratio)
if lhs > rhs + 1e-9:
viol += 1
return {"trials": trials, "violations": viol,
"max_ratio_lhs_over_rhs": round(max_ratio, 4),
"mean_ratio": round(float(np.mean(ratios)), 4),
"p95_ratio": round(float(np.percentile(ratios, 95)), 4),
"holds": viol == 0}
def check_thm56(trials=5000, perts=20):
"""Scores r_j(omega) = <a_j, omega>, so |r_j(w)-r_j(w')| <= ||a_j|| ||w-w'||;
Lipschitz L_s = max_j ||a_j||."""
inside_invariant = 0
inside_total = 0
outside_changed = 0
outside_total = 0
for _ in range(trials):
n = rng.integers(6, 40)
p = rng.integers(4, 16) # embedding dim of omega
B = rng.integers(1, n)
a = rng.normal(size=(n, p))
Ls = float(np.linalg.norm(a, axis=1).max())
omega = rng.normal(size=p)
r = a @ omega
base = topB_set(r, B)
margin = topB_margin(r, B)
if margin <= 1e-6:
continue
radius = margin / (2 * Ls)
for _ in range(perts):
direction = rng.normal(size=p)
direction /= (np.linalg.norm(direction) + 1e-12)
# inside the ball
eps_in = rng.uniform(0, radius)
r_in = a @ (omega + eps_in * direction)
inside_total += 1
if topB_set(r_in, B) == base:
inside_invariant += 1
# outside the ball (2x..6x radius) — set MAY change (tightness)
eps_out = rng.uniform(2 * radius, 6 * radius)
r_out = a @ (omega + eps_out * direction)
outside_total += 1
if topB_set(r_out, B) != base:
outside_changed += 1
return {"trials": trials,
"inside_ball_invariance_rate": round(inside_invariant / max(inside_total, 1), 6),
"inside_ball_checks": inside_total,
"outside_ball_change_rate": round(outside_changed / max(outside_total, 1), 4),
"outside_ball_checks": outside_total,
"holds": inside_invariant == inside_total}
def main():
out = {"theorem_5_4": check_thm54(), "theorem_5_6": check_thm56()}
print(json.dumps(out, indent=2))
os.makedirs("outputs", exist_ok=True)
with open("outputs/claim4_theorems.json", "w") as f:
json.dump(out, f, indent=2)
print("\nThm 5.4 holds:", out["theorem_5_4"]["holds"],
"| Thm 5.6 holds:", out["theorem_5_6"]["holds"])
if __name__ == "__main__":
main()

Xet Storage Details

Size:
4.25 kB
·
Xet hash:
e0565b9abe0b24540b9df4bb1966d08504e13bd0b582efc458728fc1e423849a

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.