File size: 18,027 Bytes
f18df08
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
"""Sentence Transformers module for fusion-embedding-2.

A thin adapter that exposes the released fusion-embedding model through the
Sentence Transformers multimodal ``encode`` API (text, image, audio, video).
All embedding math runs through the ``fusion_embedding`` package's own
``UnifiedEmbedder`` (the model's native loading path), so vectors produced here
are identical to ``fusion_embedding.UnifiedEmbedder.from_pretrained(...)``:

* text: chat-template instruction, EOS pooling, text-side whitening (fp32);
* image: the frozen base's native vision path (no whitening, no adapters);
* video: the released video preprocessing over the frozen base's video path;
* audio: soxr resampling to 16 kHz, Whisper-style mel, trained resampler, and
  the frozen decoder with ONLY the audio adapter gate open.

Every vector is L2-normalized at the full interop dimension (2048). Shorter
Matryoshka rungs: pass ``truncate_dim=<rung>`` together with
``normalize_embeddings=True`` to ``encode`` (truncate-then-renormalize equals
the native MRL readout).

Requirements (beyond sentence-transformers>=5.5.1):

    pip install "fusion-embedding[sense]>=0.3.0" torchvision

The ``sense`` extra pulls the audio decode/resample stack (soundfile, librosa);
transformers itself ships with sentence-transformers. torchvision is required to
load the model at all, not only for video, because the base processor builds a
video processor during construction; this applies to the native loader too.
Embedding an image needs Pillow, which torchvision carries. Embedding a video by
file path additionally requires torchcodec, which in turn needs FFmpeg.

Supported inputs per item (one modality per item; the model has no fused
multi-modality input):

* ``str`` — text, or a local image/audio/video file path (auto-detected);
* ``PIL.Image.Image`` or an HxWxC uint8 array — image;
* ``{"audio": {"array": waveform, "sampling_rate": sr}}`` (or the inner dict
  directly) — audio; a bare 1-D array is rejected because the sampling rate
  would be unknown;
* a ``[T, C, H, W]`` uint8 frame tensor/array — video.
"""

from __future__ import annotations

import os
from typing import Any, Optional

import torch

try:
    from sentence_transformers.base.modality import infer_modality
    from sentence_transformers.base.modules.input_module import InputModule
except ImportError as exc:  # pragma: no cover - version guard
    raise ImportError(
        "The fusion-embedding-2 Sentence Transformers integration requires "
        "sentence-transformers>=5.5.1 (multimodal encode). "
        "Upgrade with: pip install -U sentence-transformers"
    ) from exc

try:
    from fusion_embedding.config import INSTRUCTION_REGISTRY
    from fusion_embedding.model import last_token_pool
    from fusion_embedding.unified import UnifiedEmbedder, _chat
except ImportError as exc:  # pragma: no cover - dependency guard
    raise ImportError(
        "The fusion-embedding-2 Sentence Transformers integration needs the "
        "fusion-embedding package for the model implementation. Install it "
        "with: pip install 'fusion-embedding[sense]>=0.3.0'"
    ) from exc

_MIN_ST_VERSION = (5, 5, 1)


def _require_min_st_version() -> None:
    """Fail with the real reason on Sentence Transformers older than 5.5.1.

    Single-key modality dicts such as {"audio": {"array": ..., "sampling_rate": ...}}
    are classified as a tuple by infer_modality before 5.5.1, so the encode call is
    rejected by Sentence Transformers itself with a message claiming the modality is
    unsupported. The import guard above cannot catch that: sentence_transformers.base
    imports cleanly on 5.4.x.
    """
    import sentence_transformers

    raw = getattr(sentence_transformers, "__version__", "0")
    parts = []
    for chunk in raw.split(".")[:3]:
        digits = "".join(c for c in chunk if c.isdigit())
        parts.append(int(digits) if digits else 0)
    while len(parts) < 3:
        parts.append(0)
    if tuple(parts) < _MIN_ST_VERSION:
        raise ImportError(
            "The fusion-embedding-2 Sentence Transformers integration requires "
            f"sentence-transformers>=5.5.1, found {raw}. Earlier versions reject "
            "single-key modality dicts such as "
            '{"audio": {"array": ..., "sampling_rate": ...}} before this module is '
            "reached. Upgrade with: pip install -U 'sentence-transformers>=5.5.1'"
        )


CKPT_FILENAME = "fusion-embedding-2-2b-preview.pt"


class FusionEmbedding2Module(InputModule):
    """Single Sentence Transformers module wrapping the full fusion-embedding-2
    encoder (all modalities plus the canonical readout, so no separate Pooling
    or Normalize module is needed: ``forward`` emits ``sentence_embedding``
    directly)."""

    config_file_name = "sentence_bert_config.json"
    config_keys = ["ckpt_filename", "max_seq_length"]
    save_in_root = True

    def __init__(
        self,
        model_name_or_path: Optional[str] = None,
        ckpt_filename: str = CKPT_FILENAME,
        max_seq_length: int = 512,
        revision: Optional[str] = None,
        token: "bool | str | None" = None,
        cache_folder: Optional[str] = None,
        local_files_only: bool = False,
        model_kwargs: Optional[dict] = None,
        embedder: Optional[UnifiedEmbedder] = None,
        **kwargs,
    ) -> None:
        super().__init__()
        _require_min_st_version()
        self.ckpt_filename = ckpt_filename
        self.max_seq_length = max_seq_length

        if embedder is None:
            if model_name_or_path is None:
                raise ValueError("model_name_or_path is required (or pass embedder=)")
            model_kwargs = dict(model_kwargs or {})
            dtype = model_kwargs.pop("torch_dtype", model_kwargs.pop("dtype", torch.bfloat16))
            if isinstance(dtype, str):
                dtype = getattr(torch, dtype)
            device = model_kwargs.pop(
                "device", "cuda" if torch.cuda.is_available() else "cpu"
            )
            ckpt_path = self.load_file_path(
                model_name_or_path,
                filename=ckpt_filename,
                token=token,
                cache_folder=cache_folder,
                revision=revision,
                local_files_only=local_files_only,
            )
            if ckpt_path is None:
                raise FileNotFoundError(
                    f"checkpoint {ckpt_filename!r} not found in {model_name_or_path!r}"
                )
            embedder = UnifiedEmbedder.from_pretrained(ckpt_path, device=device, dtype=dtype)

        self._emb = embedder
        # Wire the video seam the UnifiedEmbedder anticipates: the released video
        # preprocessing (fusion_embedding.multimodal) over the frozen base.
        self._emb._video_pooler = self._video_pooled

        # Register the underlying torch modules so Sentence Transformers device
        # management (`model.to(device)`) moves the whole stack.
        self.fusion_model = embedder.model
        if embedder.full is not None:
            self.base = embedder.full
        if embedder.tok is not None:
            self.tokenizer = embedder.tok

    # ------------------------------------------------------------------ loading
    @classmethod
    def load(
        cls,
        model_name_or_path: str,
        subfolder: str = "",
        token: "bool | str | None" = None,
        cache_folder: Optional[str] = None,
        revision: Optional[str] = None,
        local_files_only: bool = False,
        trust_remote_code: bool = False,
        model_kwargs: Optional[dict] = None,
        processor_kwargs: Optional[dict] = None,
        config_kwargs: Optional[dict] = None,
        backend: str = "torch",
        **kwargs,
    ) -> "FusionEmbedding2Module":
        if backend != "torch":
            raise ValueError(
                f"fusion-embedding-2 only supports the torch backend, got {backend!r}"
            )
        config = cls.load_config(
            model_name_or_path,
            subfolder=subfolder,
            token=token,
            cache_folder=cache_folder,
            revision=revision,
            local_files_only=local_files_only,
        )
        config.pop("model_name_or_path", None)
        if config_kwargs:
            config.update(config_kwargs)
        return cls(
            model_name_or_path,
            revision=revision,
            token=token,
            cache_folder=cache_folder,
            local_files_only=local_files_only,
            model_kwargs=model_kwargs,
            **config,
        )

    # -------------------------------------------------------------- ST contract
    @property
    def modalities(self) -> list:
        return ["text", "image", "audio", "video"]

    def get_embedding_dimension(self) -> int:
        return int(self._emb.contract.dim)

    def save(self, output_path: str, *args, safe_serialization: bool = True, **kwargs) -> None:
        # Configuration only: the 2B weights live in the model repository's
        # checkpoint file and are not duplicated by Sentence Transformers saves.
        self.save_config(output_path)

    # ------------------------------------------------------------- input parsing
    def preprocess(self, inputs: list, prompt: Optional[str] = None, **kwargs) -> dict:
        items = []
        for item in inputs:
            modality = infer_modality(item, supported_modalities=self.modalities)
            if isinstance(modality, tuple):
                raise ValueError(
                    "fusion-embedding-2 embeds one modality per input item; "
                    f"got a combined input with {modality}. Encode each modality "
                    "separately (the shared space makes the vectors comparable)."
                )
            if isinstance(item, dict) and set(item.keys()) == {modality}:
                item = item[modality]
            items.append((modality, self._parse(modality, item)))
        return {"fusion_inputs": items, "fusion_prompt": prompt}

    def _parse(self, modality: str, item: Any) -> Any:
        if modality == "text":
            return item
        if modality == "image":
            return self._parse_image(item)
        if modality == "audio":
            return self._parse_audio(item)
        if modality == "video":
            return self._parse_video(item)
        raise ValueError(f"unsupported modality {modality!r}")

    @staticmethod
    def _parse_image(item):
        import numpy as np

        # Guarded so transformers' trust_remote_code import check does not make
        # Pillow a load-time requirement: dynamic_module_utils.get_imports skips
        # ast.Try blocks, and an unguarded import here is otherwise treated as
        # mandatory at construction even for text-only use.
        try:
            from PIL import Image
        except ImportError as exc:  # pragma: no cover - optional dependency
            raise ImportError(
                "embedding an image requires Pillow (pip install pillow)"
            ) from exc

        if isinstance(item, Image.Image):
            return item
        if isinstance(item, str):
            if item.startswith(("http://", "https://", "data:")):
                raise ValueError(
                    "image URLs / data URIs are not supported; download the file "
                    "and pass a local path or a PIL image"
                )
            return item  # local path; decoded by the native path (PIL)
        if isinstance(item, torch.Tensor):
            item = item.cpu().numpy()
        if isinstance(item, np.ndarray):
            if item.ndim == 3 and item.shape[0] in (1, 3, 4) and item.shape[-1] not in (1, 3, 4):
                item = np.transpose(item, (1, 2, 0))  # CHW -> HWC
            if item.ndim != 3 or item.shape[-1] not in (1, 3, 4):
                raise ValueError(f"expected an HxWxC image array, got shape {item.shape}")
            if item.dtype != np.uint8:
                item = np.clip(item, 0, 255).astype(np.uint8)
            return Image.fromarray(item.squeeze(-1) if item.shape[-1] == 1 else item)
        raise ValueError(f"unsupported image input type {type(item).__name__}")

    @staticmethod
    def _parse_audio(item):
        """Return (payload, sampling_rate_or_None); paths carry their own rate."""
        import numpy as np

        if isinstance(item, str):
            if item.startswith(("http://", "https://")):
                raise ValueError(
                    "audio URLs are not supported; download the file and pass a "
                    "local path or {'array': ..., 'sampling_rate': ...}"
                )
            return (item, None)
        if isinstance(item, dict):
            if "array" not in item or "sampling_rate" not in item:
                raise ValueError(
                    "audio dicts must have the form "
                    "{'array': waveform, 'sampling_rate': sr}"
                )
            array, sr = item["array"], int(item["sampling_rate"])
        else:
            try:  # torchcodec AudioDecoder (optional dependency)
                from torchcodec.decoders import AudioDecoder
            except ImportError:
                AudioDecoder = None
            if AudioDecoder is not None and isinstance(item, AudioDecoder):
                samples = item.get_all_samples()
                return (samples.data.mean(dim=0).cpu().numpy(), int(samples.sample_rate))
            raise ValueError(
                "a bare audio array has no sampling rate; pass "
                "{'audio': {'array': waveform, 'sampling_rate': sr}} instead"
            )
        if isinstance(array, torch.Tensor):
            array = array.cpu().numpy()
        array = np.asarray(array)
        if array.ndim == 2 and array.shape[0] < array.shape[1]:
            array = array.T  # (channels, samples) -> (samples, channels)
        if array.ndim > 2:
            raise ValueError(f"expected a 1-D or 2-D waveform, got shape {array.shape}")
        return (array.astype(np.float32, copy=False), sr)

    @staticmethod
    def _parse_video(item):
        import numpy as np

        if isinstance(item, str):
            if item.startswith(("http://", "https://")):
                raise ValueError(
                    "video URLs are not supported; download the file and pass a "
                    "local path or a [T, C, H, W] frame tensor"
                )
            return item  # local path; decoded natively (torchcodec, 1 fps, <=64 frames)
        if isinstance(item, dict):
            # {"array": frames, "video_metadata": ...}: the released frame-tensor
            # path derives its own metadata, so user metadata is not consumed.
            item = item["array"]
        if isinstance(item, np.ndarray):
            item = torch.from_numpy(np.ascontiguousarray(item))
        if isinstance(item, torch.Tensor):
            if item.ndim == 5 and item.shape[0] == 1:
                item = item.squeeze(0)
            if item.ndim == 4 and item.shape[-1] in (1, 3) and item.shape[1] not in (1, 3):
                item = item.permute(0, 3, 1, 2)  # THWC -> TCHW
            if item.ndim != 4:
                raise ValueError(f"expected a [T, C, H, W] frame tensor, got shape {list(item.shape)}")
            return item
        raise ValueError(f"unsupported video input type {type(item).__name__}")

    # ------------------------------------------------------------------ forward
    def forward(self, features: dict, **kwargs) -> dict:
        self._sync_device()
        prompt = features.get("fusion_prompt")
        vectors = []
        for modality, payload in features["fusion_inputs"]:
            if modality == "text":
                vectors.append(self._emb.embed_text(payload, instruction=prompt or None))
            elif modality == "image":
                vectors.append(self._emb.embed_image(payload))
            elif modality == "audio":
                array, sr = payload
                vectors.append(self._emb.embed_audio(array, sr=sr))
            elif modality == "video":
                vectors.append(self._emb.embed_video(payload))
            else:  # pragma: no cover - guarded in preprocess
                raise ValueError(f"unsupported modality {modality!r}")
        features["sentence_embedding"] = torch.stack(vectors)
        return features

    def _sync_device(self) -> None:
        """Follow Sentence Transformers device moves (`model.to(...)`)."""
        param = next(self.parameters(), None)
        if param is not None:
            self._emb.device = param.device

    # ----------------------------------------------------------- native video path
    @torch.no_grad()
    def _video_pooled(self, video, fps, max_frames) -> torch.Tensor:
        """The released fusion-embedding video path: reference-exact frame
        preprocessing (fusion_embedding.multimodal) -> frozen base's video
        forward -> EOS pooling. Runs with every adapter gate closed."""
        from fusion_embedding.config import VIDEO_USER_CONTENT
        from fusion_embedding.multimodal import _v_prepare, _v_resize_video

        emb = self._emb
        gate = getattr(emb.model, "_adapter_gate", None)
        if gate is not None and gate.active:
            raise RuntimeError("adapter gate is open during a video embed")
        if emb.full is None or emb.proc is None:
            raise RuntimeError("video embedding needs the real processor + base")
        frames, metadata = _v_prepare(video, fps, max_frames)
        frames = _v_resize_video(frames)
        text = _chat(INSTRUCTION_REGISTRY["doc"], VIDEO_USER_CONTENT)
        inputs = emb.proc(
            text=[text],
            videos=[frames],
            video_metadata=[metadata],
            do_resize=False,
            do_sample_frames=False,
            return_tensors="pt",
        ).to(emb.device)
        hidden = emb.full(**inputs).last_hidden_state
        return last_token_pool(hidden, inputs["attention_mask"])