"""InternVideo2-Stage2 1B (Wang et al., ECCV 2024, arXiv 2403.15377) — the only true VIDEO-native text channel: stage-2 trains video-text contrastively WITH temporal modeling, which the frame-pooled image channels (PE, SigLIP2, X-CLIP pooling) structurally lack. Distilled small variants are reported much weaker on retrieval — 1B or nothing. One 512-d aligned vector per episode (4-frame clip, the f4 checkpoint's native temporal extent), table iv2_vectors. Weights: ziyjiang/InternVideo2-1B (fp16 re-serialization of the official OpenGVLab stage2 checkpoint, towers only); modeling code vendored from VLM2Vec into models/iv2_stage2_1b by scripts/get_iv2.py (three MPS patches: flash-attn import guarded, LayerScale gamma naming so the checkpoint's 80 layer scales actually load, bert config resolved next to the file).""" from __future__ import annotations import numpy as np _S = {} MDIR = "models/iv2_stage2_1b" # InternVideo2's own demo preprocessing (frames2tensor): ImageNet # stats, 224 square, [B,T,C,H,W] V_MEAN = np.array([0.485, 0.456, 0.406], np.float32) V_STD = np.array([0.229, 0.224, 0.225], np.float32) def load_model(): if "model" not in _S: from transformers import AutoModel from .device import pick, strip_vision dev, dtype = pick() import transformers kw = dict(trust_remote_code=True, torch_dtype=dtype) if int(transformers.__version__.split(".")[0]) < 5: # the 8GB-container path; transformers>=5 meta-device init # breaks this custom port's from_pretrained, and the fp16 # towers load fine without it there kw["low_cpu_mem_usage"] = True m = AutoModel.from_pretrained(MDIR, **kw).to(dev).eval() m = strip_vision(m, "vision_encoder") m._config.device = dev # get_txt_feat routes tokens here _S["model"], _S["dev"], _S["dtype"] = m, dev, dtype return _S["model"], _S["dev"] def text_vec(text): cache = _S.setdefault("tcache", {}) if text in cache: return cache[text] m, _ = load_model() v = m.get_txt_feat(text).float().cpu().numpy().reshape(-1) if len(cache) > 256: cache.clear() cache[text] = v return v def set_num_frames(n): """Re-fit the vision encoder's temporal position embeddings to n frames. OPT-IN and global to the loaded model. The checkpoint ships a 4-frame clip: pos_embed is (1, 1 + 4*256, C), and feeding 8 frames raises "size of tensor a (2049) must match tensor b (1025)". That is a shape, not a capability - the patch embedding is Conv3d with a kernel of (1,14,14), so tubelet size is 1 and T frames simply need T temporal positions. The repo's own `interpolate_pos_embed(orig_t_size=4)` does exactly this, but only while loading a checkpoint. FOUR embeddings need it, not one: pos_embed AND clip_pos_embed both carry the video-length grid (the image variants are separate and untouched). Interpolating only the first fails deeper in the forward, at the CLIP-alignment branch. Sanity after interpolation: cos(4-frame, 8-frame) on the same span is 0.999, i.e. the representation is preserved rather than rebuilt. """ import torch m, _ = load_model() ve = m.vision_encoder old = int(ve.num_frames) if old == n: return def _interp(p): cls, rest = p[:, :1, :], p[:, 1:, :] C = rest.shape[-1] L = rest.shape[1] // old r = rest.view(1, old, L, C).permute(0, 3, 2, 1).float() r = torch.nn.functional.interpolate( r, size=(L, n), mode="bilinear", align_corners=False) r = r.permute(0, 3, 2, 1).reshape(1, n * L, C).to(p.dtype) return torch.nn.Parameter(torch.cat([cls, r], 1), requires_grad=False) L = (ve.pos_embed.shape[1] - 1) // old for name in ("pos_embed", "clip_pos_embed"): if hasattr(ve, name): setattr(ve, name, _interp(getattr(ve, name).data)) ve.num_frames = n ve.patch_embed.num_patches = n * L def clip_vec(frames_hwc): """Aligned 512-d vector for a clip (HWC uint8 RGB). Length must match the encoder's current num_frames - 4 by default, or whatever set_num_frames() last fitted.""" import cv2 import torch m, dev = load_model() fs = [cv2.resize(f, (224, 224)) for f in frames_hwc] x = (np.stack(fs).astype(np.float32) / 255.0 - V_MEAN) / V_STD px = torch.from_numpy(x).permute(0, 3, 1, 2)[None].to( dev, _S["dtype"]) return m.get_vid_feat(px).float().cpu().numpy().reshape(-1) def iv2_lookup(store, text): from .embeddings import _vec_table tbl, vecs = _vec_table(store, "iv2_vectors") key = {} for r, (s, a) in enumerate(zip( tbl.column("stream").to_pylist(), (int(v) for v in tbl.column("ts").to_pylist()))): key[(str(s), a)] = r sc = np.asarray(vecs) @ text_vec(text) def lookup(s, a, b): r = key.get((str(s), a)) return float(sc[r]) if r is not None else float("nan") return lookup, None