#!/usr/bin/env python """Verifica end-to-end dei port V-JEPA 2.1 pubblicati su HuggingFace. Copre le tre cose che la test suite spedita nei repo NON verifica sugli artefatti effettivamente pubblicati: A) i valori di config.json corrispondono al costruttore ufficiale (`src/hub/backbones.py::_make_vjepa2_1_model`); B) il checkpoint pubblicato si carica senza chiavi mancanti, inattese o con shape sbagliata (un parametro orfano resta inizializzato a caso e HF lo segnala solo con un warning); C) ogni tensore del safetensors pubblicato ha un'origine nel checkpoint ufficiale di Meta ed e' identico bit a bit; D) il forward dell'encoder a 384 sui pesi veri coincide con quello del reference; E) il forward del PREDICTOR sui pesi veri coincide con quello del reference (mai testato: i mask token addestrati non sono zero, quindi il test spedito con pesi random e' degenere); F) il gap di precisione ridotta misurato senza corrompere il modello. Uso su Colab ------------ !pip -q install "transformers>=4.57" safetensors huggingface_hub !pip -q install timm einops # solo per i check D/E !git clone -q https://github.com/facebookresearch/vjepa2.git !python verify_vjepa21_port.py --repo apiantonio/vjepa2.1-vit-base-384 \ --vjepa2-repo ./vjepa2 --checks ABCDEF RAM richiesta (i check C/D/E tengono in memoria port + reference): base ~2 GB | large ~4 GB | giant ~13 GB | gigantic ~20 GB Su Colab free (12.7 GB) girano base e large; per giant/gigantic serve una runtime High-RAM, oppure si eseguono solo A e B. """ from __future__ import annotations import argparse import copy import json import math import os import sys import urllib.request import torch # Meta ha lasciato VJEPA_BASE_URL puntato a http://localhost:8300 nel main # corrente di facebookresearch/vjepa2 (la riga vera e' commentata sopra), quindi # torch.hub.load(...) fallisce. Scarichiamo il checkpoint direttamente. OFFICIAL_URL = "https://dl.fbaipublicfiles.com/vjepa2" GREEN, RED, YELLOW, RESET = "\033[32m", "\033[31m", "\033[33m", "\033[0m" def ok(msg): print(f"{GREEN} PASS{RESET} {msg}") def fail(msg): print(f"{RED} FAIL{RESET} {msg}") FAILURES.append(msg) def warn(msg): print(f"{YELLOW} WARN{RESET} {msg}") FAILURES: list[str] = [] # --------------------------------------------------------------------------- # A) config.json vs costruttore ufficiale # --------------------------------------------------------------------------- # Derivato da src/hub/backbones.py::_make_vjepa2_1_model + le factory in # app/vjepa_2_1/models/vision_transformer.py. NON toccare senza rileggere il # reference: e' la specifica contro cui si verifica. COMMON = dict( patch_size=16, crop_size=384, tubelet_size=2, frames_per_clip=64, # num_frames=64 in_chans=3, img_temporal_dim_size=1, interpolate_rope=True, modality_embedding=True, hidden_act="gelu", # use_silu=False qkv_bias=True, n_registers=0, has_cls_first=False, layer_norm_eps=1e-6, drop_path_rate=0.0, num_pooler_layers=3, # +1 cross-attention block = num_probe_blocks: 4 num_pooler_heads=16, # classifier.num_heads: 16 in every configs/eval_2_1 file pred_hidden_size=384, # predictor_embed_dim pred_num_attention_heads=12, # num_heads=12 nel predictor pred_mlp_ratio=4.0, pred_num_mask_tokens=8, # predictor_num_mask_tokens pred_zero_init_mask_tokens=True, pred_return_all_tokens=True, # return_all_tokens=True ) EXPECTED = { "apiantonio/vjepa2.1-vit-base-384": dict( COMMON, hidden_size=768, num_hidden_layers=12, num_attention_heads=12, mlp_ratio=4.0, n_output_distillation=1, pred_num_hidden_layers=12, pred_teacher_embed_dim=1664, _ckpt="vjepa2_1_vitb_dist_vitG_384.pt", _key="ema_encoder", _arch="vit_base", ), "apiantonio/vjepa2.1-vit-large-384": dict( COMMON, hidden_size=1024, num_hidden_layers=24, num_attention_heads=16, mlp_ratio=4.0, n_output_distillation=1, pred_num_hidden_layers=12, pred_teacher_embed_dim=1664, _ckpt="vjepa2_1_vitl_dist_vitG_384.pt", _key="ema_encoder", _arch="vit_large", ), "apiantonio/vjepa2.1-vit-giant-384": dict( COMMON, hidden_size=1408, num_hidden_layers=40, num_attention_heads=22, mlp_ratio=48 / 11, n_output_distillation=4, pred_num_hidden_layers=24, pred_teacher_embed_dim=None, _ckpt="vjepa2_1_vitg_384.pt", _key="target_encoder", _arch="vit_giant_xformers", ), "apiantonio/vjepa2.1-vit-gigantic-384": dict( COMMON, hidden_size=1664, num_hidden_layers=48, num_attention_heads=26, mlp_ratio=64 / 13, n_output_distillation=4, pred_num_hidden_layers=24, pred_teacher_embed_dim=None, _ckpt="vjepa2_1_vitG_384.pt", _key="target_encoder", _arch="vit_gigantic_xformers", ), } def check_config(repo): print(f"\n[A] config.json vs costruttore ufficiale — {repo}") from transformers import AutoConfig cfg = AutoConfig.from_pretrained(repo, trust_remote_code=True) spec = EXPECTED[repo] bad = [] for key, want in spec.items(): if key.startswith("_"): continue got = getattr(cfg, key, "") same = math.isclose(got, want, rel_tol=0, abs_tol=0) if isinstance(want, float) else got == want if not same: bad.append(f"{key}: atteso {want!r}, trovato {got!r}") if bad: for b in bad: fail(b) else: ok(f"{len([k for k in spec if not k.startswith(chr(95))])} campi coincidono") # le proprieta' derivate devono coincidere con la mappa del reference hier = {12: [2, 5, 8, 11], 24: [5, 11, 17, 23], 40: [9, 19, 29, 39], 48: [11, 23, 37, 47]} if cfg.encoder_hierarchical_layers != hier[cfg.num_hidden_layers]: fail(f"encoder_hierarchical_layers {cfg.encoder_hierarchical_layers}") else: ok(f"encoder_hierarchical_layers = {cfg.encoder_hierarchical_layers}") # dimensione della proiezione del predictor n_hier = len(cfg.predictor_hierarchical_layers) out = (cfg.pred_teacher_embed_dim // n_hier) if cfg.pred_teacher_embed_dim else cfg.hidden_size ok(f"predictor proj out_dim = {n_hier * out} (n_hier={n_hier})") # MLP: int(dim * ratio) deve dare esattamente il valore del reference for name, d, r in (("encoder", cfg.hidden_size, cfg.mlp_ratio), ("predictor", cfg.pred_hidden_size, cfg.pred_mlp_ratio)): h = int(d * r) exact = int(d * (48 / 11)) if abs(r - 48 / 11) < 1e-12 else ( int(d * (64 / 13)) if abs(r - 64 / 13) < 1e-12 else int(d * r)) if h != exact: fail(f"{name} mlp hidden {h} != {exact} (round-trip JSON del mlp_ratio)") else: ok(f"{name} mlp hidden = {h}") return cfg # --------------------------------------------------------------------------- # B) nessun parametro orfano al caricamento # --------------------------------------------------------------------------- def check_loading(repo, dtype=torch.float32): print(f"\n[B] caricamento senza chiavi orfane — {repo}") from transformers import AutoModel, AutoModelForVideoClassification model, info = AutoModel.from_pretrained( repo, trust_remote_code=True, dtype=dtype, output_loading_info=True ) for name in ("missing_keys", "unexpected_keys", "mismatched_keys"): v = info.get(name) or [] (ok if not v else fail)(f"AutoModel {name}: {len(v)}" + (f" -> {v[:6]}" if v else "")) n = sum(p.numel() for p in model.parameters()) ok(f"parametri totali: {n:,}") # la testa di classificazione deve ereditare i pesi dell'encoder: # se base_model_prefix e' sbagliato, missing_keys esplode e il backbone # riparte da zero senza che nulla fallisca. clf, cinfo = AutoModelForVideoClassification.from_pretrained( repo, trust_remote_code=True, dtype=dtype, num_labels=2, output_loading_info=True ) missing = [k for k in (cinfo.get("missing_keys") or []) if not k.startswith(("pooler.", "classifier."))] (ok if not missing else fail)( f"ForVideoClassification: solo pooler/classifier reinizializzati" + (f", ma anche {missing[:6]}" if missing else "") ) for k, v in model.encoder.state_dict().items(): if not torch.equal(v, clf.vjepa21.encoder.state_dict()[k]): fail(f"peso encoder diverso dopo il wrapper: {k}") break else: ok("i pesi dell'encoder sopravvivono al wrapper di classificazione") del clf return model # --------------------------------------------------------------------------- # mappa reference -> port (identica a quella usata dalla conversione) # --------------------------------------------------------------------------- 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 in ("pos_embed",): continue # non usato: il modello usa RoPE else: warn(f"chiave encoder di reference 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 in ("predictor_pos_embed",): continue else: warn(f"chiave predictor di reference non mappata: {k}") return out def download_official(repo, cache="."): name = EXPECTED[repo]["_ckpt"] path = os.path.join(cache, name) if not os.path.exists(path): print(f" scarico {name} ...") urllib.request.urlretrieve(f"{OFFICIAL_URL}/{name}", path) return path def load_official(repo, cache="."): path = download_official(repo, cache) raw = torch.load(path, map_location="cpu", weights_only=False) clean = lambda sd: {k.replace("module.", "").replace("backbone.", ""): v for k, v in sd.items()} return clean(raw[EXPECTED[repo]["_key"]]), clean(raw["predictor"]) # --------------------------------------------------------------------------- # C) provenienza bit a bit di ogni tensore pubblicato # --------------------------------------------------------------------------- def check_provenance(repo, cache="."): print(f"\n[C] provenienza dei pesi pubblicati — {repo}") from huggingface_hub import hf_hub_download from safetensors.torch import load_file cfg = EXPECTED[repo] enc_sd, pred_sd = load_official(repo, cache) expected = reference_to_port(enc_sd, pred_sd, cfg["hidden_size"], cfg["pred_hidden_size"]) published = load_file(hf_hub_download(repo, "model.safetensors")) orphans = sorted(set(published) - set(expected)) unused = sorted(set(expected) - set(published)) (ok if not orphans else fail)( f"tensori pubblicati senza origine nel checkpoint: {len(orphans)}" + (f" -> {orphans[:8]}" if orphans else "") ) if unused: warn(f"tensori del reference non pubblicati: {len(unused)} -> {unused[:8]}") worst, worst_key = 0.0, None for k in sorted(set(published) & set(expected)): a, b = published[k].float(), expected[k].float() if a.shape != b.shape: fail(f"shape diversa per {k}: {tuple(a.shape)} vs {tuple(b.shape)}") continue d = (a - b).abs().max().item() if d > worst: worst, worst_key = d, k (ok if worst == 0.0 else fail)( f"max|Δ| su {len(set(published) & set(expected))} tensori = {worst:.3e}" + (f" (peggiore: {worst_key})" if worst else "") ) del published, expected, enc_sd, pred_sd # --------------------------------------------------------------------------- # D/E) parita' del forward sui pesi veri, encoder e predictor # --------------------------------------------------------------------------- def build_reference(repo, vjepa2_repo, cache="."): 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 = EXPECTED[repo] enc = vit.__dict__[spec["_arch"]]( patch_size=16, img_size=(384, 384), num_frames=64, tubelet_size=2, use_sdpa=False, uniform_power=False, use_rope=True, img_temporal_dim_size=1, interpolate_rope=True, modality_embedding=True, n_output_distillation=spec["n_output_distillation"], ).eval() # NOTE: `VisionTransformerPredictor.__init__` has no `use_sdpa` parameter — # it would be swallowed by `**kwargs` — so the reference predictor blocks # always run SDPA while the port runs eager. That is why the predictor # tolerance below is 1e-3 rather than exact. pred = vit_predictor( img_size=(384, 384), patch_size=16, use_mask_tokens=True, embed_dim=spec["hidden_size"], predictor_embed_dim=384, teacher_embed_dim=spec["pred_teacher_embed_dim"], num_frames=64, tubelet_size=2, depth=spec["pred_num_hidden_layers"], num_heads=12, num_mask_tokens=8, use_rope=True, uniform_power=False, use_silu=False, wide_silu=True, n_output_distillation=spec["n_output_distillation"], return_all_tokens=True, img_temporal_dim_size=1, modality_embedding=True, zero_init_mask_tokens=True, interpolate_rope=True, ).eval() enc_sd, pred_sd = load_official(repo, cache) enc.load_state_dict(enc_sd, strict=True) pred.load_state_dict(pred_sd, strict=True) ok("encoder e predictor di reference caricati con strict=True") return enc, pred @torch.no_grad() def check_forward_parity(repo, vjepa2_repo, port, frames=16, cache="."): print(f"\n[D] parita' del forward encoder a 384, pesi pubblicati — {repo}") ref_enc, ref_pred = build_reference(repo, vjepa2_repo, cache) torch.manual_seed(0) x = torch.randn(1, 3, frames, 384, 384) saved_impl = getattr(port.config, "_attn_implementation", "sdpa") port.config._attn_implementation = "eager" # il reference encoder usa use_sdpa=False a = ref_enc(x) b = port(pixel_values_videos=x, skip_predictor=True).last_hidden_state d = (a - b).abs().max().item() (ok if d < 1e-4 else fail)(f"T={frames}: max|Δ| = {d:.3e} (tokens={a.shape[1]})") # [E] predictor sui pesi VERI: i mask token addestrati non sono zero, quindi # questo esercita davvero il percorso che il test spedito non copre. print(f"\n[E] parita' del forward predictor, pesi pubblicati — {repo}") mt = torch.stack([m.flatten() for m in ref_pred.mask_tokens]).abs().max().item() (warn if mt == 0 else ok)(f"norma max dei mask token del checkpoint = {mt:.3e}" + (" (zero: il test resta degenere)" if mt == 0 else "")) z = ref_enc(x, training=True) if EXPECTED[repo]["n_output_distillation"] > 1 else a N = z.shape[1] ctx = torch.arange(0, N // 2).unsqueeze(0) tgt = torch.arange(N // 2, N).unsqueeze(0) from importlib import import_module apply_masks = import_module(type(port).__module__).apply_masks rp, rc = ref_pred(apply_masks(z, [ctx]), [ctx], [tgt], mod="video") # `VisionTransformerPredictor.__init__` non ha `use_sdpa` (finisce in **kwargs), # quindi il predictor del reference gira SEMPRE in SDPA. Appaiamo il kernel, # altrimenti si misura la differenza eager-vs-SDPA e non la parita' del port. port.config._attn_implementation = "sdpa" got = port.predictor(z, [ctx], [tgt], mode="video") port.config._attn_implementation = saved_impl def _r(u, v): return ((u - v).abs().mean() / u.abs().mean()).item() def _c(u, v): return torch.nn.functional.cosine_similarity( u.flatten(0, 1), v.flatten(0, 1)).min().item() pr, pc = _r(rp, got.last_hidden_state), _c(rp, got.last_hidden_state) cr, cc = _r(rc, got.context_hidden_state), _c(rc, got.context_hidden_state) peak = (rp - got.last_hidden_state).abs().max().item() # L'assert e' su errore relativo e cosine similarity: il massimo assoluto e' # preso su milioni di elementi e non significa nulla senza la scala delle # attivazioni. (ok if pr < 1e-4 and pc > 0.9999 and cr < 1e-4 and cc > 0.9999 else fail)( f"target rel = {pr:.3e} cos = {pc:.6f} | context rel = {cr:.3e} cos = {cc:.6f}" f" | target peak = {peak:.3e}" ) del ref_enc, ref_pred # --------------------------------------------------------------------------- # F) precisione ridotta senza corrompere il modello # --------------------------------------------------------------------------- @torch.no_grad() def check_precision(port): print("\n[F] gap di precisione ridotta (su copia, il modello non viene alterato)") device = "cuda" if torch.cuda.is_available() else "cpu" torch.manual_seed(1) x = torch.randn(1, 3, 4, port.config.crop_size, port.config.crop_size, device=device) base = port.to(device) ref = base(pixel_values_videos=x, skip_predictor=True).last_hidden_state.float() saved = getattr(base.config, "_attn_implementation", "sdpa") 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"): low = copy.deepcopy(base).to(dtype) # <- la copia e' il punto low.config._attn_implementation = impl out = low(pixel_values_videos=x.to(dtype), skip_predictor=True).last_hidden_state finite = bool(torch.isfinite(out).all()) got = out.float() label = f"{str(dtype).split('.')[-1]}/{impl}" # V-JEPA 2.1 e' addestrato in bfloat16 (use_bfloat16: true nei config di # eval), che ha il range di esponente del fp32: le attivazioni possono # uscire dal range fp16. E' un limite del checkpoint, non del port. if not finite: warn(f"{label:>16}: NaN/inf, overflow fp16 nel kernel {impl}") else: rel = ((ref - got).abs().mean() / ref.abs().mean()).item() cos = torch.nn.functional.cosine_similarity( ref.flatten(0, 1), got.flatten(0, 1)).min().item() (ok if cos > 0.99 else fail)( f"{label:>16}: rel = {rel:.3e} min cos = {cos:.6f}") del low base.config._attn_implementation = saved a = base(pixel_values_videos=x, skip_predictor=True).last_hidden_state b = base(pixel_values_videos=x, skip_predictor=True).last_hidden_state (ok if torch.equal(a, b) else fail)("due forward identici sono bit-identici") # --------------------------------------------------------------------------- # G) DoRA: il merge deve essere esatto quanto quello di LoRA # --------------------------------------------------------------------------- @torch.no_grad() def check_dora(repo): print("\n[G] merge di DoRA (modello piccolo, la proprieta' e' algebrica)") try: from peft import LoraConfig, get_peft_model except ImportError: warn("peft non installato, salto") return 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 cls = AutoModelForVideoClassification.from_config(small, trust_remote_code=True) for use_dora in (False, True): torch.manual_seed(0) model = copy.deepcopy(cls).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)}") # --------------------------------------------------------------------------- def main(): ap = argparse.ArgumentParser() ap.add_argument("--repo", required=True, choices=sorted(EXPECTED)) ap.add_argument("--vjepa2-repo", default="./vjepa2") ap.add_argument("--cache", default=".") ap.add_argument("--frames", type=int, default=16) ap.add_argument("--checks", default="ABCDEFG") args = ap.parse_args() print(f"torch {torch.__version__} | cuda {torch.cuda.is_available()}") port = None if "A" in args.checks: check_config(args.repo) if set("BDEFG") & set(args.checks): port = check_loading(args.repo) if "B" in args.checks else None if port is None: from transformers import AutoModel port = AutoModel.from_pretrained(args.repo, trust_remote_code=True).eval() port.eval() if "C" in args.checks: check_provenance(args.repo, args.cache) if set("DE") & set(args.checks): check_forward_parity(args.repo, args.vjepa2_repo, port, args.frames, args.cache) if "F" in args.checks: check_precision(port) if "G" in args.checks: check_dora(args.repo) print("\n" + "=" * 70) if FAILURES: print(f"{RED}{len(FAILURES)} controlli falliti{RESET}") for f in FAILURES: print(" -", f) sys.exit(1) print(f"{GREEN}tutti i controlli eseguiti sono passati{RESET}") if __name__ == "__main__": main()