SudharshanR
ElideDB query by example: no text, no model at query time
a1dd5ba
Raw
History Blame Contribute Delete
5.11 kB
"""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