File size: 5,111 Bytes
a1dd5ba
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
"""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