vjepa2.1-vit-base-384 / verify_big_models.py
apiantonio's picture
Fix transformers 4.x/5.x compat, implement output_hidden_states/attentions and out_layers, fix hierarchical predictor input, add video processor
21dd249 verified
Raw
History Blame Contribute Delete
23.9 kB
#!/usr/bin/env python
"""Check C/D/E per i checkpoint grandi (ViT-g, ViT-G) senza andare in OOM.
`verify_vjepa21_port.py` tiene in RAM contemporaneamente il .pt di Meta, la
mappa dei tensori attesi, il safetensors pubblicato, il modello di reference e
il port. Su base e large ci sta; su giant (1.07 G parametri) e gigantic (1.90 G)
no. Qui le stesse tre verifiche sono riorganizzate cosi':
[C] provenienza — il .pt viene aperto in mmap (le pagine restano su disco) e
il safetensors pubblicato viene letto UN TENSORE ALLA VOLTA
con `safe_open`. Gestisce anche i checkpoint shardati:
gigantic e' distribuito in 4 file piu' un index.json, e
`hf_hub_download(repo, "model.safetensors")` fallisce.
Picco: qualche centinaio di MB.
[D] encoder — il reference e il port NON sono mai vivi insieme. Fase 1:
[E] predictor costruisci il reference, calcola, scrivi gli output su
disco, libera. Fase 2: carica il port, calcola, confronta.
Picco: un solo modello alla volta.
[F] precisione — niente `copy.deepcopy` del modello gia' su GPU: il modello
a precisione ridotta viene ricaricato da disco direttamente
nel dtype voluto. Con gigantic il deepcopy chiedeva 7.1 GB
di VRAM oltre ai 7.1 gia' occupati, contro i 14.6 di una T4.
[G] merge PEFT — gira su un modello piccolo derivato dalla config reale: il
merge esatto e' una proprieta' algebrica, non dei pesi.
Entrambi i lati girano in SDPA. Il reference encoder accetta `use_sdpa=True`, e
il predictor del reference usa SDPA comunque (`use_sdpa` non e' un suo
parametro, finisce in **kwargs). Oltre a essere il confronto corretto a kernel
appaiato, evita di materializzare la matrice di attenzione: a 4608 token con 22
teste sarebbero 1.9 GB per layer con il kernel eager.
Uso:
!python verify_big_models.py --repo apiantonio/vjepa2.1-vit-giant-384 \
--vjepa2-repo ./vjepa2 --checks CDE
Su Colab free (12.7 GB di RAM) girano entrambi. Serve spazio su disco per il
.pt di Meta (~4 GB per giant, ~8 GB per gigantic) piu' il safetensors.
"""
from __future__ import annotations
import argparse
import gc
import json
import os
import resource
import sys
import urllib.request
import torch
OFFICIAL_URL = "https://dl.fbaipublicfiles.com/vjepa2"
SPEC = {
"apiantonio/vjepa2.1-vit-base-384": dict(
ckpt="vjepa2_1_vitb_dist_vitG_384.pt", key="ema_encoder", arch="vit_base",
hidden=768, n_distill=1, pred_depth=12, teacher=1664),
"apiantonio/vjepa2.1-vit-large-384": dict(
ckpt="vjepa2_1_vitl_dist_vitG_384.pt", key="ema_encoder", arch="vit_large",
hidden=1024, n_distill=1, pred_depth=12, teacher=1664),
"apiantonio/vjepa2.1-vit-giant-384": dict(
ckpt="vjepa2_1_vitg_384.pt", key="target_encoder", arch="vit_giant_xformers",
hidden=1408, n_distill=4, pred_depth=24, teacher=None),
"apiantonio/vjepa2.1-vit-gigantic-384": dict(
ckpt="vjepa2_1_vitG_384.pt", key="target_encoder", arch="vit_gigantic_xformers",
hidden=1664, n_distill=4, pred_depth=24, teacher=None),
}
G, R, Y, N = "\033[32m", "\033[31m", "\033[33m", "\033[0m"
FAILURES: list[str] = []
PRED_DIM = 384
def ok(m):
print(f"{G} PASS{N} {m}")
def fail(m):
print(f"{R} FAIL{N} {m}")
FAILURES.append(m)
def warn(m):
print(f"{Y} WARN{N} {m}")
def peak_ram_gb() -> float:
return resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / (1024 * 1024)
def report_ram(tag=""):
print(f" [picco RAM {peak_ram_gb():.2f} GB{(' — ' + tag) if tag else ''}]")
# ---------------------------------------------------------------------------
# accesso frugale ai due checkpoint
# ---------------------------------------------------------------------------
def official_path(repo, cache="."):
name = SPEC[repo]["ckpt"]
path = os.path.join(cache, name)
if not os.path.exists(path):
print(f" scarico {name} (una volta sola) ...")
urllib.request.urlretrieve(f"{OFFICIAL_URL}/{name}", path)
print(f" {name}: {os.path.getsize(path) / 2**30:.2f} GB su disco")
return path
def load_official_mmap(repo, cache="."):
"""Apre il .pt in mmap e scarta subito tutto cio' che non serve.
I checkpoint di training di Meta contengono anche l'encoder non-EMA e lo
stato dell'ottimizzatore: possono pesare 3-4 volte il modello.
"""
path = official_path(repo, cache)
try:
raw = torch.load(path, map_location="cpu", mmap=True, weights_only=False)
except (RuntimeError, TypeError) as e:
warn(f"mmap non disponibile ({type(e).__name__}), carico normalmente")
raw = torch.load(path, map_location="cpu", weights_only=False)
key = SPEC[repo]["key"]
keep = {key, "predictor"}
dropped = [k for k in list(raw.keys()) if k not in keep]
print(f" chiavi nel .pt: {sorted(raw.keys())}")
for k in dropped:
del raw[k]
gc.collect()
clean = lambda sd: {k.replace("module.", "").replace("backbone.", ""): v
for k, v in sd.items()}
return clean(raw[key]), clean(raw["predictor"])
def published_tensors(repo):
"""Genera (nome, tensore) leggendo un tensore alla volta.
Gestisce sia il file unico sia i checkpoint shardati (gigantic ha 4 shard
piu' `model.safetensors.index.json`).
"""
from huggingface_hub import hf_hub_download
from safetensors import safe_open
try:
files = [hf_hub_download(repo, "model.safetensors")]
except Exception:
idx = hf_hub_download(repo, "model.safetensors.index.json")
with open(idx) as fh:
shards = sorted(set(json.load(fh)["weight_map"].values()))
print(f" checkpoint shardato in {len(shards)} file")
files = [hf_hub_download(repo, s) for s in shards]
for f in files:
with safe_open(f, framework="pt", device="cpu") as h:
for name in h.keys():
yield name, h.get_tensor(name)
# ---------------------------------------------------------------------------
# mappa reference -> port (identica a verify_vjepa21_port.py)
# ---------------------------------------------------------------------------
def _map_block(prefix, idx, sub, tensor, hidden, out):
if sub.startswith("attn.qkv."):
kind = sub.rsplit(".", 1)[-1]
q, k, v = tensor.split(hidden, dim=0)
out[f"{prefix}.layer.{idx}.attention.query.{kind}"] = q
out[f"{prefix}.layer.{idx}.attention.key.{kind}"] = k
out[f"{prefix}.layer.{idx}.attention.value.{kind}"] = v
elif sub.startswith("attn.proj."):
out[f"{prefix}.layer.{idx}.attention.proj." + sub.rsplit(".", 1)[-1]] = tensor
else:
out[f"{prefix}.layer.{idx}.{sub}"] = tensor
def reference_to_port(enc_sd, pred_sd, hidden, pred_hidden):
out = {}
for k, v in enc_sd.items():
if k in ("img_mod_embed", "video_mod_embed"):
out[f"encoder.embeddings.{k}"] = v
elif k.startswith("patch_embed_img."):
out["encoder.embeddings.patch_embeddings_img." + k[len("patch_embed_img."):]] = v
elif k.startswith("patch_embed."):
out["encoder.embeddings.patch_embeddings." + k[len("patch_embed."):]] = v
elif k.startswith("norms_block."):
out["encoder." + k] = v
elif k.startswith("blocks."):
idx, sub = k[len("blocks."):].split(".", 1)
_map_block("encoder", idx, sub, v, hidden, out)
elif k == "pos_embed":
continue
else:
warn(f"chiave encoder non mappata: {k}")
for k, v in pred_sd.items():
if k in ("img_mod_embed", "video_mod_embed"):
out[f"predictor.embeddings.{k}"] = v
elif k.startswith("predictor_embed."):
out["predictor.embeddings.predictor_embed." + k[len("predictor_embed."):]] = v
elif k.startswith("mask_tokens."):
out["predictor.embeddings." + k] = v
elif k.startswith("predictor_norm."):
out["predictor.layernorm." + k[len("predictor_norm."):]] = v
elif k.startswith("predictor_proj_context."):
out["predictor.proj_context." + k[len("predictor_proj_context."):]] = v
elif k.startswith("predictor_proj."):
out["predictor.proj." + k[len("predictor_proj."):]] = v
elif k.startswith("predictor_blocks."):
idx, sub = k[len("predictor_blocks."):].split(".", 1)
_map_block("predictor", idx, sub, v, pred_hidden, out)
elif k == "predictor_pos_embed":
continue
else:
warn(f"chiave predictor non mappata: {k}")
return out
# ---------------------------------------------------------------------------
# [C] provenienza in streaming
# ---------------------------------------------------------------------------
def check_provenance(repo, cache="."):
print(f"\n[C] provenienza dei pesi pubblicati — {repo}")
spec = SPEC[repo]
enc_sd, pred_sd = load_official_mmap(repo, cache)
expected = reference_to_port(enc_sd, pred_sd, spec["hidden"], PRED_DIM)
print(f" tensori attesi dalla conversione: {len(expected)}")
seen, orphans, worst, worst_key, mismatched = 0, [], 0.0, None, []
for name, tensor in published_tensors(repo):
seen += 1
ref = expected.pop(name, None)
if ref is None:
orphans.append(name)
continue
if tuple(ref.shape) != tuple(tensor.shape):
mismatched.append(f"{name}: {tuple(tensor.shape)} vs {tuple(ref.shape)}")
continue
d = (tensor.float() - ref.float()).abs().max().item()
if d > worst:
worst, worst_key = d, name
del tensor, ref
(ok if not orphans else fail)(
f"tensori pubblicati senza origine nel checkpoint: {len(orphans)}"
+ (f" -> {sorted(orphans)[:6]}" if orphans else ""))
(ok if not mismatched else fail)(
f"tensori con shape diversa: {len(mismatched)}"
+ (f" -> {mismatched[:4]}" if mismatched else ""))
if expected:
warn(f"tensori del reference non pubblicati: {len(expected)} "
f"-> {sorted(expected)[:6]}")
(ok if worst == 0.0 else fail)(
f"max|Δ| su {seen - len(orphans)} tensori = {worst:.3e}"
+ (f" (peggiore: {worst_key})" if worst else ""))
report_ram("dopo C")
del expected, enc_sd, pred_sd
gc.collect()
# ---------------------------------------------------------------------------
# [D] + [E] parita' del forward, reference e port mai vivi insieme
# ---------------------------------------------------------------------------
@torch.no_grad()
def phase_reference(repo, vjepa2_repo, frames, device, work, cache="."):
print(f"\n[D/E] fase 1 — reference (il port non e' ancora caricato)")
sys.path.insert(0, os.path.abspath(vjepa2_repo))
from app.vjepa_2_1.models import vision_transformer as vit
from app.vjepa_2_1.models.predictor import vit_predictor
spec = SPEC[repo]
enc_sd, pred_sd = load_official_mmap(repo, cache)
encoder = vit.__dict__[spec["arch"]](
patch_size=16, img_size=(384, 384), num_frames=64, tubelet_size=2,
use_sdpa=True, uniform_power=False, use_rope=True, img_temporal_dim_size=1,
interpolate_rope=True, modality_embedding=True,
n_output_distillation=spec["n_distill"],
).eval()
encoder.load_state_dict(enc_sd, strict=True)
del enc_sd
gc.collect()
ok("encoder di reference caricato con strict=True")
torch.manual_seed(0)
x = torch.randn(1, 3, frames, 384, 384)
encoder = encoder.to(device)
# `forward(training=True)` restituisce i livelli concatenati; l'ultima fetta e'
# `norms_block[-1]` applicata all'ultimo layer, cioe' esattamente il
# `last_hidden_state`. Vale sia per n_distill=1 sia per n_distill=4.
z = encoder(x.to(device), training=True).cpu()
a = z[..., -spec["hidden"]:].contiguous()
del encoder
gc.collect()
if device == "cuda":
torch.cuda.empty_cache()
print(f" encoder: {tuple(a.shape)} token, hierarchical {tuple(z.shape)}")
report_ram("dopo l'encoder di reference")
predictor = vit_predictor(
img_size=(384, 384), patch_size=16, use_mask_tokens=True,
embed_dim=spec["hidden"], predictor_embed_dim=PRED_DIM,
teacher_embed_dim=spec["teacher"], num_frames=64, tubelet_size=2,
depth=spec["pred_depth"], num_heads=12, num_mask_tokens=8,
use_rope=True, uniform_power=False, use_silu=False, wide_silu=True,
n_output_distillation=spec["n_distill"], return_all_tokens=True,
img_temporal_dim_size=1, modality_embedding=True, zero_init_mask_tokens=True,
interpolate_rope=True,
).eval()
predictor.load_state_dict(pred_sd, strict=True)
del pred_sd
gc.collect()
ok("predictor di reference caricato con strict=True")
mt = max(m.abs().max().item() for m in predictor.mask_tokens)
(warn if mt == 0 else ok)(
f"norma max dei mask token del checkpoint = {mt:.3e}"
+ (" (zero: il confronto sul mask token resta degenere)" if mt == 0 else ""))
n_tokens = z.shape[1]
ctx = torch.arange(0, n_tokens // 2).unsqueeze(0)
tgt = torch.arange(n_tokens // 2, n_tokens).unsqueeze(0)
predictor = predictor.to(device)
idx = ctx.unsqueeze(-1).expand(-1, -1, z.size(-1)).to(device)
ctx_tokens = torch.gather(z.to(device), 1, idx)
rp, rc = predictor(ctx_tokens, [ctx.to(device)], [tgt.to(device)], mod="video")
rp, rc = rp.cpu(), rc.cpu()
del predictor, ctx_tokens, idx
gc.collect()
if device == "cuda":
torch.cuda.empty_cache()
report_ram("dopo il predictor di reference")
blob = os.path.join(work, "reference_outputs.pt")
torch.save({"x": x, "a": a, "z": z, "rp": rp, "rc": rc, "ctx": ctx, "tgt": tgt}, blob)
print(f" output di reference scritti in {blob} "
f"({os.path.getsize(blob) / 2**30:.2f} GB)")
return blob
@torch.no_grad()
def phase_port(repo, blob, device):
print(f"\n[D/E] fase 2 — port (il reference e' stato liberato)")
from transformers import AutoModel
cached = torch.load(blob, map_location="cpu", weights_only=False)
port = AutoModel.from_pretrained(repo, trust_remote_code=True).eval()
port.config._attn_implementation = "sdpa"
port = port.to(device)
report_ram("port caricato")
def rel(u, v):
return ((u - v).abs().mean() / u.abs().mean()).item()
def cos(u, v):
return torch.nn.functional.cosine_similarity(
u.flatten(0, 1), v.flatten(0, 1)).min().item()
# --- D ---
out = port(pixel_values_videos=cached["x"].to(device), skip_predictor=True,
return_hierarchical=True)
b = out.last_hidden_state.cpu()
zh = out.hierarchical_hidden_state.cpu()
del out
if device == "cuda":
torch.cuda.empty_cache()
a = cached["a"]
d, r, c = (a - b).abs().max().item(), rel(a, b), cos(a, b)
(ok if (d == 0.0 or (r < 1e-6 and c > 0.999999)) else fail)(
f"[D] encoder: max|Δ| = {d:.3e} rel = {r:.3e} cos = {c:.6f} "
f"(tokens={a.shape[1]})")
dz = (cached["z"] - zh).abs().max().item()
(ok if dz == 0.0 else fail)(
f"[D] hierarchical: max|Δ| = {dz:.3e} (dim={zh.shape[-1]})")
del b, zh
gc.collect()
# --- E ---
got = port.predictor(cached["z"].to(device),
[cached["ctx"].to(device)], [cached["tgt"].to(device)],
mode="video")
gp = got.last_hidden_state.cpu()
gc_ = got.context_hidden_state.cpu()
del got, port
gc.collect()
if device == "cuda":
torch.cuda.empty_cache()
rp, rc = cached["rp"], cached["rc"]
pr, pc = rel(rp, gp), cos(rp, gp)
cr, cc = rel(rc, gc_), cos(rc, gc_)
peak = (rp - gp).abs().max().item()
(ok if pr < 1e-4 and pc > 0.9999 and cr < 1e-4 and cc > 0.9999 else fail)(
f"[E] predictor: target rel = {pr:.3e} cos = {pc:.6f} | "
f"context rel = {cr:.3e} cos = {cc:.6f} | max|Δ| = {peak:.3e} "
f"(out dim = {gp.shape[-1]})")
report_ram("fine")
# ---------------------------------------------------------------------------
# [F] precisione ridotta senza tenere due copie del modello in memoria
# ---------------------------------------------------------------------------
def _from_pretrained(repo, dtype=None):
"""`dtype=` su transformers 5, `torch_dtype=` su transformers 4."""
from transformers import AutoModel
kw = dict(trust_remote_code=True)
if dtype is not None:
try:
return AutoModel.from_pretrained(repo, dtype=dtype, **kw)
except TypeError:
return AutoModel.from_pretrained(repo, torch_dtype=dtype, **kw)
return AutoModel.from_pretrained(repo, **kw)
@torch.no_grad()
def check_precision_streaming(repo, device):
"""Come il check F di verify_vjepa21_port.py, ma senza `copy.deepcopy`.
Il deepcopy avviene sul modello gia' spostato su GPU: per gigantic sono
7.1 GB in fp32 piu' altri 7.1 GB per la copia, contro i 14.6 GB di una T4.
Qui il modello a precisione ridotta viene RICARICATO da disco direttamente
nel dtype voluto, quindi non c'e' mai piu' di un modello per volta. E' anche
piu' pulito del deepcopy rispetto al problema originale: ogni forward parte
da pesi freschi, quindi nessuna misura puo' contaminare la successiva.
"""
print(f"\n[F] gap di precisione ridotta — {repo}")
if device == "cuda":
torch.cuda.reset_peak_memory_stats()
model = _from_pretrained(repo).eval()
crop = model.config.crop_size
torch.manual_seed(1)
x = torch.randn(1, 3, 4, crop, crop)
model = model.to(device)
ref = model(pixel_values_videos=x.to(device), skip_predictor=True).last_hidden_state
twice = model(pixel_values_videos=x.to(device), skip_predictor=True).last_hidden_state
(ok if torch.equal(ref, twice) else fail)("due forward identici sono bit-identici")
ref = ref.float().cpu()
del model, twice
gc.collect()
if device == "cuda":
torch.cuda.empty_cache()
print(f" [picco VRAM fp32 {torch.cuda.max_memory_allocated() / 2**30:.2f} GB]")
for dtype in (torch.bfloat16, torch.float16):
if dtype is torch.float16 and device == "cpu":
warn("fp16 saltato su CPU")
continue
for impl in ("sdpa", "eager"):
if device == "cuda":
torch.cuda.reset_peak_memory_stats()
low = _from_pretrained(repo, dtype=dtype).eval()
low.config._attn_implementation = impl
low = low.to(device)
out = low(pixel_values_videos=x.to(device).to(dtype),
skip_predictor=True).last_hidden_state
finite = bool(torch.isfinite(out).all())
got = out.float().cpu()
del low, out
gc.collect()
if device == "cuda":
torch.cuda.empty_cache()
label = f"{str(dtype).split('.')[-1]}/{impl}"
if not finite:
# V-JEPA 2.1 e' addestrato in bfloat16, che ha il range di esponente
# del fp32: i logits dell'attenzione possono uscire dal range fp16, e
# il kernel eager li calcola in fp16 nativo. Limite del checkpoint.
warn(f"{label:>16}: NaN/inf, overflow fp16 nel kernel eager")
else:
r = ((ref - got).abs().mean() / ref.abs().mean()).item()
c = torch.nn.functional.cosine_similarity(
ref.flatten(0, 1), got.flatten(0, 1)).min().item()
(ok if c > 0.99 else fail)(f"{label:>16}: rel = {r:.3e} min cos = {c:.6f}")
del got
report_ram("dopo F")
# ---------------------------------------------------------------------------
# [G] merge di LoRA e DoRA — modello piccolo, la proprieta' e' algebrica
# ---------------------------------------------------------------------------
@torch.no_grad()
def check_dora(repo):
print(f"\n[G] merge di LoRA e DoRA (modello piccolo)")
try:
from peft import LoraConfig, get_peft_model
except ImportError:
warn("peft non installato, salto")
return
import copy as _copy
from transformers import AutoConfig, AutoModelForVideoClassification
cfg = AutoConfig.from_pretrained(repo, trust_remote_code=True)
small = _copy.deepcopy(cfg)
small.crop_size, small.hidden_size, small.num_attention_heads = 64, 96, 6
small.num_hidden_layers, small.pred_num_hidden_layers = 12, 12
small.pred_hidden_size, small.pred_num_attention_heads = 48, 6
small.n_output_distillation, small.pred_teacher_embed_dim = 1, 96
small.num_pooler_heads, small.num_labels = 6, 5
base = AutoModelForVideoClassification.from_config(small, trust_remote_code=True)
for use_dora in (False, True):
torch.manual_seed(0)
model = _copy.deepcopy(base).eval()
peft_model = get_peft_model(model, LoraConfig(
r=8, lora_alpha=16, lora_dropout=0.0, use_dora=use_dora,
target_modules=r".*vjepa21\.encoder\.layer\.\d+\.attention\.(query|key|value|proj)$",
modules_to_save=["classifier", "pooler"],
))
for name, p in peft_model.named_parameters():
if "lora_B" in name:
torch.nn.init.normal_(p, std=0.02)
x = torch.randn(2, 3, 4, 64, 64)
before = peft_model(pixel_values_videos=x).logits
merged = peft_model.merge_and_unload().eval()
after = merged(pixel_values_videos=x).logits
d = (before - after).abs().max().item()
leftover = [n for n, _ in merged.named_parameters() if "lora" in n.lower()]
label = "DoRA" if use_dora else "LoRA"
(ok if d < 1e-4 and not leftover else fail)(
f"{label}: max|Δ| dopo merge = {d:.3e}, tensori adapter residui = {len(leftover)}")
del model, peft_model, merged
gc.collect()
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--repo", required=True, choices=sorted(SPEC))
ap.add_argument("--vjepa2-repo", default="./vjepa2")
ap.add_argument("--cache", default=".", help="dove tenere il .pt di Meta")
ap.add_argument("--work", default="/tmp", help="dove scrivere gli output intermedi")
ap.add_argument("--frames", type=int, default=16)
ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
ap.add_argument("--checks", default="CDEFG")
ap.add_argument("--keep-blob", action="store_true")
args = ap.parse_args()
print(f"torch {torch.__version__} | device {args.device} | repo {args.repo}")
print(f"RAM iniziale in uso: {peak_ram_gb():.2f} GB")
if "C" in args.checks:
check_provenance(args.repo, args.cache)
blob = None
if set("DE") & set(args.checks):
blob = phase_reference(args.repo, args.vjepa2_repo, args.frames,
args.device, args.work, args.cache)
gc.collect()
phase_port(args.repo, blob, args.device)
if blob and not args.keep_blob:
os.remove(blob)
gc.collect()
if "F" in args.checks:
check_precision_streaming(args.repo, args.device)
gc.collect()
if "G" in args.checks:
check_dora(args.repo)
print("\n" + "=" * 70)
print(f"picco RAM del processo: {peak_ram_gb():.2f} GB")
if FAILURES:
print(f"{R}{len(FAILURES)} controlli falliti{N}")
for f in FAILURES:
print(" -", f)
sys.exit(1)
print(f"{G}tutti i controlli eseguiti sono passati{N}")
if __name__ == "__main__":
main()