# 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)