Spaces:
Running
Running
Download app.py from misukisu/embeddinggemma-2-api: direct link, hf CLI and curl.
- Browser
- Download file 4.76 kB
-
https://huggingface.co/spaces/misukisu/embeddinggemma-2-api/resolve/main/app.py
- Command line
-
hf download hf://spaces/misukisu/embeddinggemma-2-api/app.py
-
curl -L -o app.py https://huggingface.co/spaces/misukisu/embeddinggemma-2-api/resolve/main/app.py
4.76 kB
| import os | |
| import io | |
| import base64 | |
| import torch | |
| import numpy as np | |
| import av | |
| import soundfile as sf | |
| from PIL import Image | |
| from typing import Optional, List | |
| from fastapi import FastAPI, UploadFile, File, Form, HTTPException | |
| from fastapi.responses import Response | |
| from pydantic import BaseModel | |
| import orjson | |
| from transformers import AutoProcessor, AutoModel | |
| torch.set_num_threads(2) | |
| torch.set_num_interop_threads(1) | |
| os.environ["OMP_NUM_THREADS"] = "2" | |
| app = FastAPI(title="EmbeddingGemma 2 Multimodal API") | |
| MODEL_ID = os.getenv("MODEL_PATH", "google/embeddinggemma-2") | |
| processor = AutoProcessor.from_pretrained(MODEL_ID) | |
| model = AutoModel.from_pretrained(MODEL_ID, torch_dtype=torch.float32) | |
| model.eval() | |
| class ORJSONCustomResponse(Response): | |
| media_type = "application/json" | |
| def render(self, content: any) -> bytes: | |
| return orjson.dumps(content) | |
| def truncate_and_normalize(tensor: torch.Tensor, dim: int = 256) -> List[float]: | |
| vec = tensor[0, :dim] | |
| norm = torch.linalg.vector_norm(vec) | |
| normalized = vec / torch.clamp(norm, min=1e-12) | |
| return normalized.tolist() | |
| def extract_video_frames(video_bytes: bytes, max_frames: int = 16) -> List[Image.Image]: | |
| container = av.open(io.BytesIO(video_bytes)) | |
| frames = [] | |
| for frame in container.decode(video=0): | |
| frames.append(frame.to_image().convert("RGB")) | |
| if len(frames) >= 120: | |
| break | |
| container.close() | |
| if not frames: | |
| raise ValueError("Could not decode any video frames") | |
| indices = np.linspace(0, len(frames) - 1, min(len(frames), max_frames), dtype=int) | |
| return [frames[i] for i in indices] | |
| def read_audio(audio_bytes: bytes, target_sr: int = 16000) -> np.ndarray: | |
| with sf.SoundFile(io.BytesIO(audio_bytes)) as audio_file: | |
| audio = audio_file.read(dtype="float32") | |
| sr = audio_file.samplerate | |
| if audio.ndim > 1: | |
| audio = audio.mean(axis=1) | |
| if sr != target_sr: | |
| duration = len(audio) / sr | |
| new_len = int(duration * target_sr) | |
| audio = np.interp(np.linspace(0, len(audio), new_len), np.arange(len(audio)), audio) | |
| return audio | |
| class JSONMultimodalRequest(BaseModel): | |
| text: Optional[str] = None | |
| image_b64: Optional[str] = None | |
| audio_b64: Optional[str] = None | |
| task_prefix: Optional[str] = "SearchQuery" | |
| dimensions: Optional[int] = 256 | |
| def health(): | |
| return {"status": "ready", "engine": "multimodal-embeddinggemma-2"} | |
| def embed_json(payload: JSONMultimodalRequest): | |
| dim = payload.dimensions if payload.dimensions in [128, 256, 512, 768] else 768 | |
| images = None | |
| audio = None | |
| if payload.image_b64: | |
| raw_img = base64.b64decode(payload.image_b64) | |
| images = [Image.open(io.BytesIO(raw_img)).convert("RGB")] | |
| if payload.audio_b64: | |
| raw_audio = base64.b64decode(payload.audio_b64) | |
| audio = read_audio(raw_audio) | |
| text_input = payload.text | |
| if text_input and payload.task_prefix and not images and not audio: | |
| text_input = f"{payload.task_prefix}: {text_input}" | |
| inputs = processor( | |
| text=text_input if text_input else None, | |
| images=images, | |
| audio=audio, | |
| return_tensors="pt" | |
| ) | |
| with torch.inference_mode(): | |
| outputs = model(**inputs) | |
| embedding = outputs.last_hidden_state[:, 0, :] | |
| vec = truncate_and_normalize(embedding, dim) | |
| return ORJSONCustomResponse({"embedding": vec, "dimensions": dim}) | |
| def embed_file( | |
| file: UploadFile = File(...), | |
| text: Optional[str] = Form(None), | |
| task_prefix: Optional[str] = Form("SearchQuery"), | |
| dimensions: int = Form(256) | |
| ): | |
| dim = dimensions if dimensions in [128, 256, 512, 768] else 768 | |
| contents = file.file.read() | |
| content_type = file.content_type or "" | |
| images = None | |
| audio = None | |
| if content_type.startswith("image/"): | |
| images = [Image.open(io.BytesIO(contents)).convert("RGB")] | |
| elif content_type.startswith("video/"): | |
| images = extract_video_frames(contents, max_frames=8) | |
| elif content_type.startswith("audio/"): | |
| audio = read_audio(contents) | |
| else: | |
| raise HTTPException(status_code=400, detail=f"Unsupported media type: {content_type}") | |
| inputs = processor( | |
| text=text if text else None, | |
| images=images, | |
| audio=audio, | |
| return_tensors="pt" | |
| ) | |
| with torch.inference_mode(): | |
| outputs = model(**inputs) | |
| embedding = outputs.last_hidden_state[:, 0, :] | |
| vec = truncate_and_normalize(embedding, dim) | |
| return ORJSONCustomResponse({ | |
| "filename": file.filename, | |
| "type": content_type, | |
| "embedding": vec, | |
| "dimensions": dim | |
| }) | |