File size: 4,773 Bytes
6e49b2b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
# 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)