veil-pgd / ensemble /probe_towers.py
Klaus Clawd
Initial public release: VEIL-PGD v0.1
c793f45
Raw
History Blame Contribute Delete
5.06 kB
"""Empirical load + differentiable-gradient probe for candidate vision towers.
We ONLY trust a tower if it loads on this transformers and yields a real input
gradient. Prefer newest-of-family; report FAIL reasons so we can pick fallbacks.
Run on a Spark: python -m ensemble.probe_towers
"""
from __future__ import annotations
import time
import traceback
import torch
_D = "cuda"
def log(m):
print(f"[{time.strftime('%H:%M:%S')}] {m}", flush=True)
def _grad_ok(feat, x):
g = torch.autograd.grad(feat.float().pow(2).sum(), x, retain_graph=False)[0]
return float(g.float().norm().item())
def _as_feat(out):
"""Normalize any HF/transformers vision output to a (1, D) feature."""
import torch as _t
if isinstance(out, (tuple, list)):
out = out[0]
h = getattr(out, "last_hidden_state", None)
if h is None:
h = getattr(out, "pooler_output", None)
if h is None and isinstance(out, _t.Tensor):
h = out
if h.dim() == 3:
h = h.mean(dim=1)
elif h.dim() == 2 and h.shape[0] > 1:
h = h.mean(dim=0, keepdim=True)
return h
def probe_moonvit():
from transformers import AutoModel
m = AutoModel.from_pretrained("moonshotai/MoonViT-SO-400M", torch_dtype=torch.float16,
trust_remote_code=True).to(_D).eval()
for p in m.parameters():
p.requires_grad_(False)
x = torch.rand(1, 3, 448, 448, device=_D, dtype=torch.float16, requires_grad=True)
# MoonViT wants pixel_values + grid; try a few signatures
try:
out = m(x)
except Exception:
gh = torch.tensor([[448 // 14, 448 // 14]], device=_D)
out = m(x, gh)
h = out[0] if isinstance(out, (tuple, list)) else getattr(out, "last_hidden_state", out)
feat = h.reshape(1, -1) if h.dim() > 2 else h
return feat, x
def probe_internvit():
from transformers import AutoModel
m = AutoModel.from_pretrained("OpenGVLab/InternViT-300M-448px-V2_5",
torch_dtype=torch.float16, trust_remote_code=True,
attn_implementation="sdpa").to(_D).eval()
for p in m.parameters():
p.requires_grad_(False)
x = torch.rand(1, 3, 448, 448, device=_D, dtype=torch.float16, requires_grad=True)
out = m(pixel_values=x)
h = getattr(out, "last_hidden_state", None)
if h is None:
h = out[0]
return h.mean(dim=1), x
def probe_qwen35():
# Newest Qwen. Try to instantiate ONLY the vision tower from config to avoid
# loading the 9B LLM.
from transformers import AutoConfig, AutoModel
repo = "Qwen/Qwen3.5-9B"
cfg = AutoConfig.from_pretrained(repo, trust_remote_code=True)
# full model (fp16, ~18GB) then grab .visual / .model.visual
from transformers import AutoModelForImageTextToText
m = AutoModelForImageTextToText.from_pretrained(repo, torch_dtype=torch.float16,
device_map=_D, trust_remote_code=True).eval()
vis = getattr(getattr(m, "model", m), "visual", None) or getattr(m, "visual", None)
for p in m.parameters():
p.requires_grad_(False)
# Qwen vision takes flattened patches + grid_thw; build a minimal grid
ps, tps = cfg.vision_config.patch_size, cfg.vision_config.temporal_patch_size
gh = gw = 16
seq = gh * gw
dim = 3 * tps * ps * ps
x = torch.rand(seq, dim, device=_D, dtype=torch.float16, requires_grad=True)
grid = torch.tensor([[1, gh, gw]], device=_D)
out = vis(x, grid_thw=grid)
return _as_feat(out), x
def probe_pixtral_ministral3():
from transformers import AutoModelForImageTextToText
repo = "mistralai/Ministral-3-8B-Instruct-2512"
m = AutoModelForImageTextToText.from_pretrained(repo, torch_dtype=torch.float16,
device_map=_D).eval()
for p in m.parameters():
p.requires_grad_(False)
vm = getattr(getattr(m, "model", m), "vision_tower", None) or \
getattr(getattr(m, "model", m), "vision_model", None)
x = torch.rand(1, 3, 512, 512, device=_D, dtype=torch.float16, requires_grad=True)
out = vm(x)
h = getattr(out, "last_hidden_state", None)
if h is None:
h = out[0]
return h.mean(dim=1), x
PROBES = {
"MoonViT-SO-400M (Kimi K2.7)": probe_moonvit,
"InternViT-300M (InternVL3.5)": probe_internvit,
"Qwen3.5-9B .visual": probe_qwen35,
"Pixtral via Ministral-3-8B": probe_pixtral_ministral3,
}
def main():
for name, fn in PROBES.items():
t = time.time()
try:
feat, x = fn()
gl2 = _grad_ok(feat, x)
log(f"OK {name}: feat={tuple(feat.shape)} grad_l2={gl2:.3f} "
f"({time.time()-t:.0f}s) VRAM={torch.cuda.max_memory_allocated()/1e9:.1f}GB")
except Exception as e: # noqa: BLE001
log(f"FAIL {name}: {repr(e)[:200]}")
traceback.print_exc()
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
if __name__ == "__main__":
main()