#!/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()