qwenjev / scripts /vision_probe2.py
tchbcb's picture
v0.4.9: RCA evidence archived + HANDOVER §0.4
6e49b2b verified
Raw History Blame Contribute Delete
4.77 kB
# DECISIVE probe: is the vision tower alive, and where does the signal die?
# Linear code (exec scope safe). Loads the fp16 model once.
import os, sys
os.environ["HF_HUB_DISABLE_XET"] = "1"
sys.path.insert(0, "/content/qwenjev")
sys.path.insert(0, "/content/qwenjev/train")
import torch
import torch.nn.functional as F
from pathlib import Path
from qwenjev.hf_compat import load_multimodal, find_decoder_layers
from qwenjev.vision import render_grid_pil
from transformers import AutoProcessor
DEV = "cuda"
proc = AutoProcessor.from_pretrained("Qwen/Qwen3-VL-2B-Instruct")
model = load_multimodal("Qwen/Qwen3-VL-2B-Instruct", device=DEV)
model.eval()
print("model loaded", flush=True)
# drastically different frames (32x32, like a 109-cell change but bigger)
gA = [[(r * 32 + c) % 16 for c in range(32)] for r in range(32)]
gB = [[((r * 32 + c) * 7 + 3) % 16 for c in range(32)] for r in range(32)]
imgA, imgB = render_grid_pil(gA), render_grid_pil(gB)
STATE_TXT = "STATE:\n" + ("\n".join("".join(str((r * 7) % 16) for r in range(8)) for _ in range(32))) + "\n\nQUESTION:\nWhich candidate action is most likely to make progress in this game?\n\nCANDIDATE:"
def build_batch(with_image):
texts = []
for cand in ("ACTION1", "ACTION6"):
content = ([{"type": "image"}] if with_image else []) + [
{"type": "text", "text": STATE_TXT + " " + cand}]
texts.append(proc.apply_chat_template(
[{"role": "user", "content": content}],
add_generation_prompt=True, tokenize=False))
kw = {"images": [imgA, imgB]} if with_image else {}
return proc(text=texts, padding=True, return_tensors="pt", **kw).to(DEV)
batch_img = build_batch(True)
batch_no = build_batch(False)
print("img batch keys:", sorted(batch_img.keys()), flush=True)
print("vision tokens:", int((batch_img["input_ids"] == 151655).sum()), flush=True)
layers = find_decoder_layers(model)
picks = {}
def hook(_m, _inp, out):
t = out[0] if isinstance(out, tuple) else out
picks["h"] = t
h = layers[-1].register_forward_hook(hook)
with torch.no_grad():
out_img = model(**batch_img, logits_to_keep=1)
hids = picks["h"].float() # [B, T, H]
out_img.logits = None
h.remove()
print("hidden:", tuple(hids.shape), flush=True)
VID = 151655
res = {}
for i, tag in ((0, "A"), (1, "B")):
ids = batch_img["input_ids"][i]
vpos = (ids == VID).nonzero().flatten()
vh = hids[i, vpos] # [Nv, H]
last = hids[i, batch_img["attention_mask"][i].sum() - 1]
res[tag] = {"vh": vh, "last": last}
print(f"frame {tag}: vision_pos={len(vpos)} vh_norm={vh.norm(dim=-1).mean():.2f} last_norm={last.norm():.2f}", flush=True)
cos_vh_mean = F.cosine_similarity(res["A"]["vh"].mean(0), res["B"]["vh"].mean(0), dim=0).item()
cos_vh_per = sum(F.cosine_similarity(res["A"]["vh"], res["B"]["vh"], dim=-1)).item() / len(res["A"]["vh"])
cos_last = F.cosine_similarity(res["A"]["last"], res["B"]["last"], dim=0).item()
print(f"COS vision-token mean-vec A vs B : {cos_vh_mean:.5f}", flush=True)
print(f"COS vision-token per-token A vs B : {cos_vh_per:.5f}", flush=True)
print(f"COS last-position A vs B : {cos_last:.5f} (diag said ~0.9999)", flush=True)
# image vs no-image at the last position (same text)
picks.clear()
h = layers[-1].register_forward_hook(hook)
with torch.no_grad():
out_no = model(**batch_no, logits_to_keep=1)
hids_no = picks["h"].float()
out_no.logits = None
h.remove()
cos_last_noimg = F.cosine_similarity(
hids[0, batch_img["attention_mask"][0].sum() - 1],
hids_no[0, batch_no["attention_mask"][0].sum() - 1], dim=0).item()
print(f"COS last-pos WITH-image vs NO-image (frame A) : {cos_last_noimg:.5f}", flush=True)
# the tower itself: run vision encoder on both frames
vis = None
for name, mod in model.named_modules():
if name.endswith("visual") and hasattr(mod, "forward") and len(list(mod.children())) > 2:
vis = mod
print("vision tower module:", name, type(mod).__name__, flush=True)
break
if vis is not None:
with torch.no_grad():
tA = vis(batch_img["pixel_values"][0:1],
grid_thw=batch_img["image_grid_thw"][0:1])
tB = vis(batch_img["pixel_values"][1:2],
grid_thw=batch_img["image_grid_thw"][1:2])
tA = tA[-1] if isinstance(tA, tuple) else tA
tB = tB[-1] if isinstance(tB, tuple) else tB
tA, tB = tA.float()[0], tB.float()[0]
print(f"tower out: {tuple(tA.shape)} normA={tA.norm():.1f}", flush=True)
print(f"COS tower-output frame A vs B : {F.cosine_similarity(tA, tB, dim=-1).mean().item():.5f}", flush=True)
print(f"L2 tower-output frame A vs B : {(tA - tB).norm().item():.3f} (vs normA {tA.norm().item():.1f})", flush=True)
print("PROBE2_DONE", flush=True)