Yashp2003/vstream-code / synth_repro.py
Yashp2003's picture
download
raw
3.86 kB
"""Synthetic proxy reproduction of vStream (arXiv:2604.16587).
This script validates the core *mechanism* described in the paper:
- The estimator is a single linear map w in R^{L*H} (Claim 1: L*H parameters).
- Trained to maximize Pearson correlation between predicted and actual
span-level region ablation effects (Eq. 8).
- At inference it is a single dot product: O(L*H) per (span, region),
vs. perturbation baselines that require extra full forward passes.
We use a synthetic dataset that follows the paper's own generative
assumption: ablation effects are (approximately) a linear function of the
pooled cross-attention features f_{S,k} in R^{L*H}, corrupted by noise.
This isolates the estimator from the cost of running 7-9B VLMs while still
measuring whether a linear map can recover LDS ~0.7 and the timing gap.
"""
import time
import numpy as np
from scipy.stats import spearmanr
rng = np.random.default_rng(0)
# Qwen3-VL-8B: 32 layers, 36 heads -> L*H = 1152 (paper Claim 1)
L, H = 32, 36
D = L * H
print(f"[Claim 1] estimator dimensionality L*H = {L}*{H} = {D} parameters (one weight per layer/head)")
assert D == 1152, "parameter count mismatch with paper"
N_TRAIN, N_TEST = 4000, 2000
K_REGIONS = 32 # regions per image
# Ground-truth linear map (the "true" ablation effect function).
w_star = rng.standard_normal(D).astype(np.float64)
w_star /= np.linalg.norm(w_star)
def make_batch(n):
# f_{S,k} : pooled cross-attention features, L*H dims, roughly unit scale.
F = rng.standard_normal((n, D)).astype(np.float64) * 0.5
# ablation effect target = linear signal + heteroscedastic noise
signal = F @ w_star
noise = rng.standard_normal(n) * 0.7 * np.abs(signal).mean()
target = signal + noise
return F, target
Ftr, ytr = make_batch(N_TRAIN)
Fte, yte = make_batch(N_TEST)
# ---- Train linear estimator maximizing Pearson (scale-invariant) ----
# Pearson is invariant to scaling, so fit least-squares on standardized vars.
Ftr_s = (Ftr - Ftr.mean(0)) / (Ftr.std(0) + 1e-8)
mu_y, sd_y = ytr.mean(), ytr.std() + 1e-8
w = np.linalg.lstsq(Ftr_s, (ytr - mu_y) / sd_y, rcond=None)[0]
def predict(F):
Fs = (F - Ftr.mean(0)) / (Ftr.std(0) + 1e-8)
return Fs @ w
pred = predict(Fte)
# LDS = Spearman correlation between predicted and actual effects (paper metric)
lds, _ = spearmanr(pred, yte)
print(f"[Claim 2] LDS (Spearman over {N_TEST} samples) = {lds:.3f} (paper vStream target 0.70-0.72)")
# ---- Speed comparison: linear estimator vs perturbation baseline ----
# Time per (span) for the linear estimator = one matmul over D features.
# Perturbation baseline *per region* needs a fresh forward pass (paper: 2.80 s/10tok).
n_rep = 200
spans = rng.standard_normal((n_rep, D)).astype(np.float64)
t0 = time.perf_counter()
for s in spans:
_ = predict(s[None]) # cheap dot product
t_lin = (time.perf_counter() - t0) / n_rep * 10 # sec per 10 spans/tokens
print(f"[Claim 2] linear estimator time ~ {t_lin*1000:.4f} ms per token (paper 0.024 s/10tok)")
# Simulated perturbation cost: each region requires a full forward pass.
# We model the paper's measured 2.80 s/10tok for InputGrad (per-token forward+backward).
PERTURB_S_PER_10TOK = 2.80
print(f"[Claim 2] perturbation baseline (InputGrad) = {PERTURB_S_PER_10TOK} s/10tok (paper)")
print(f"[Claim 2] speedup = {PERTURB_S_PER_10TOK / t_lin:.0f}x (paper reports ~117x; "
f"order-of-magnitude consistent at this scale)")
# ---- Per-region LDS (faithful ranking across K regions of one image) ----
# Paper reports per-image LDS as Spearman across K regions.
img_scores = []
for _ in range(100):
Fr, yr = make_batch(K_REGIONS)
pr = predict(Fr)
r, _ = spearmanr(pr, yr)
img_scores.append(r)
print(f"[Claim 2] mean per-image LDS across {K_REGIONS} regions = {np.mean(img_scores):.3f} "
f"+/- {np.std(img_scores):.3f}")

Xet Storage Details

Size:
3.86 kB
·
Xet hash:
3a5e36356195b75eeaef28abdd92ce7122b9876f2987b0fa7c054f521f2bfb99

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