OmniRet / omniret /embedding.py
chuonghm's picture
Publish OmniRet checkpoint-2220
03f1aba verified
Raw
History Blame Contribute Delete
7.88 kB
from __future__ import annotations
import json
import tempfile
import wave
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from urllib.parse import urlparse
from urllib.request import urlopen
import imageio_ffmpeg
import numpy as np
import torch
import torch.nn.functional as F
from huggingface_hub import snapshot_download
from PIL import Image
from safetensors.torch import load_model as load_safetensors_model
from .config import OmniRetModelConfig
from .modeling import OmniRetModel
_MEDIA_FIELDS = ("audio", "image", "video")
_INPUT_FIELDS = {"instruction", "text", "audio", "image", "video", "max_frames"}
@dataclass(frozen=True)
class _Input:
instruction: str
text: str
modality: str | None
media: str | Path | None
max_frames: int
def _validate_input(value: dict[str, Any], index: int) -> _Input:
if not isinstance(value, dict):
raise TypeError(f"input {index} must be a dictionary")
unknown = sorted(set(value) - _INPUT_FIELDS)
if unknown:
raise ValueError(f"input {index} has unknown keys: {unknown}")
present = [field for field in _MEDIA_FIELDS if value.get(field) is not None]
if len(present) > 1:
raise ValueError(f"input {index} has multiple media fields")
text = str(value.get("text") or "").strip()
if not text and not present:
raise ValueError(f"input {index} has no text or media")
media = value.get(present[0]) if present else None
if media is not None and not isinstance(media, (str, Path)):
raise TypeError(f"input {index} {present[0]} must be a path or HTTP(S) URL")
max_frames = value.get("max_frames", 8)
if not isinstance(max_frames, int) or not 1 <= max_frames <= 8:
raise ValueError(f"input {index} max_frames must be between 1 and 8")
return _Input(
instruction=str(value.get("instruction") or "").strip(),
text=text,
modality=present[0] if present else None,
media=media,
max_frames=max_frames,
)
def _validate_inputs(inputs: list[dict[str, Any]]) -> list[_Input]:
if not isinstance(inputs, list):
raise TypeError("inputs must be a list")
if not inputs:
raise ValueError("inputs must not be empty")
return [_validate_input(value, index) for index, value in enumerate(inputs)]
def _format_input(value: _Input) -> str:
pieces = []
if value.modality:
pieces.append(f"{value.modality.title()}: <{value.modality}>")
if value.text:
pieces.append(value.text)
body = "\n".join(pieces)
return f"Instruct: {value.instruction}\nQuery:\n{body}" if value.instruction else body
@contextmanager
def _local_path(value: str | Path, index: int, field: str):
text = str(value)
if text.startswith(("http://", "https://")):
suffix = Path(urlparse(text).path).suffix
try:
with urlopen(text, timeout=30) as response, tempfile.NamedTemporaryFile(suffix=suffix) as handle:
payload = response.read(512 * 1024 * 1024 + 1)
if len(payload) > 512 * 1024 * 1024:
raise ValueError(f"input {index} {field} download exceeds 512 MiB")
handle.write(payload)
handle.flush()
yield Path(handle.name)
except OSError as error:
raise ValueError(f"input {index} could not download {field}: {error}") from error
return
path = Path(value)
if not path.is_file():
raise FileNotFoundError(f"input {index} {field} file does not exist: {path}")
yield path
def _load_wav(path: Path, index: int) -> np.ndarray:
with wave.open(str(path)) as handle:
if (handle.getnchannels(), handle.getsampwidth(), handle.getframerate()) != (1, 2, 16000):
raise ValueError(f"input {index} expected mono 16-bit 16 kHz PCM WAV")
return np.frombuffer(handle.readframes(handle.getnframes()), dtype=np.int16).astype(np.float32) / 32768
def _load_video(path: Path, max_frames: int, index: int) -> list[Image.Image]:
reader = imageio_ffmpeg.read_frames(path)
try:
metadata = next(reader)
finally:
reader.close()
total = max(1, int(float(metadata["fps"]) * float(metadata["duration"])))
indices = np.linspace(0, total - 1, min(max_frames, total)).round().astype(int).tolist()
select = "+".join(f"eq(n\\,{frame})" for frame in indices)
reader = imageio_ffmpeg.read_frames(path, output_params=["-vf", f"select={select}", "-vsync", "0"])
try:
metadata = next(reader)
frames = [Image.frombytes("RGB", metadata["size"], frame) for frame in reader]
finally:
reader.close()
if not frames:
raise ValueError(f"input {index} video contains no decodable frames")
return frames
def _load_media(value: str | Path, modality: str, max_frames: int, index: int):
with _local_path(value, index, modality) as path:
try:
if modality == "audio":
return _load_wav(path, index)
if modality == "video":
return _load_video(path, max_frames, index)
with Image.open(path) as image:
return image.convert("RGB")
except (OSError, RuntimeError, ValueError, wave.Error) as error:
if isinstance(error, ValueError) and str(error).startswith(f"input {index}"):
raise
raise ValueError(f"input {index} could not decode {modality}: {error}") from error
def _model_dir(model_name_or_path: str | Path) -> Path:
path = Path(model_name_or_path)
return path if path.is_dir() else Path(snapshot_download(repo_id=str(model_name_or_path)))
class OmniRetEmbedder:
def __init__(
self,
model_name_or_path: str | Path,
torch_dtype: torch.dtype | None = None,
attn_implementation: str | None = None,
device: str | torch.device | None = None,
) -> None:
model_dir = _model_dir(model_name_or_path)
metadata = json.loads((model_dir / "config.json").read_text())
bases = {
name: snapshot_download(repo_id=value["repo_id"], revision=value["revision"])
for name, value in metadata["base_models"].items()
}
config = OmniRetModelConfig(
text_model_name=bases["text"],
vision_model_name=bases["vision"],
audio_model_name=bases["audio"],
attn_implementation=attn_implementation,
)
self.device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu"))
dtype = torch_dtype or (torch.bfloat16 if self.device.type == "cuda" else torch.float32)
self.model = OmniRetModel.from_pretrained(config).to(dtype=dtype)
load_safetensors_model(self.model, str(model_dir / "model.safetensors"), strict=True)
self.model = self.model.to(device=self.device).eval()
@torch.inference_mode()
def process(self, inputs: list[dict[str, Any]]) -> torch.Tensor:
values = _validate_inputs(inputs)
texts = [_format_input(value) for value in values]
modalities = [value.modality or "image" for value in values]
media = [
None if value.media is None else _load_media(value.media, modalities[index], value.max_frames, index)
for index, value in enumerate(values)
]
raw = self.model.encode_raw_media_batch(media, modalities) if any(item is not None for item in media) else None
tokens, media_mask = raw if isinstance(raw, tuple) else (raw, None)
embeddings, _ = self.model.encode_batch(
texts,
tokens,
modalities,
media_mask,
exclude_instruction_prefix=any(value.instruction for value in values),
)
return F.normalize(embeddings.float(), dim=-1)