Video Classification
Transformers
Safetensors
vjepa21
feature-extraction
video
vjepa
vjepa2
v-jepa-2.1
self-supervised
world-model
custom_code
Instructions to use apiantonio/vjepa2.1-vit-gigantic-384 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use apiantonio/vjepa2.1-vit-gigantic-384 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("video-classification", model="apiantonio/vjepa2.1-vit-gigantic-384", trust_remote_code=True)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("apiantonio/vjepa2.1-vit-gigantic-384", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Fix transformers 4.x/5.x compat, implement output_hidden_states/attentions and out_layers, fix hierarchical predictor input, add video processor
28a7a26 verified | #!/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, "<assente>") | |
| 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 | |
| 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 | |
| # --------------------------------------------------------------------------- | |
| 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 | |
| # --------------------------------------------------------------------------- | |
| 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() |