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-base-384 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use apiantonio/vjepa2.1-vit-base-384 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("video-classification", model="apiantonio/vjepa2.1-vit-base-384", trust_remote_code=True)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("apiantonio/vjepa2.1-vit-base-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
21dd249 verified | #!/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 | |
| # --------------------------------------------------------------------------- | |
| 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 | |
| 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) | |
| 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 | |
| # --------------------------------------------------------------------------- | |
| 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() |