"""EmbeddingGemma 2 multimodal API (API-only; no web frontend). Run: uvicorn app:app --host 0.0.0.0 --port 7860 --workers 1 """ from __future__ import annotations import base64 import binascii import io import os import tempfile import threading from contextlib import asynccontextmanager from pathlib import Path from typing import Annotated, Literal import numpy as np import orjson import soundfile as sf import torch from fastapi import FastAPI, File, Form, HTTPException, UploadFile from fastapi.concurrency import run_in_threadpool from fastapi.responses import PlainTextResponse, Response from PIL import Image, UnidentifiedImageError from pydantic import BaseModel, Field, model_validator from transformers import AutoModel, AutoProcessor # Keep the CPU footprint manageable on small HF Spaces. NUM_THREADS = int(os.getenv("TORCH_THREADS", "2")) torch.set_num_threads(NUM_THREADS) os.environ.setdefault("OMP_NUM_THREADS", str(NUM_THREADS)) MODEL_ID = os.getenv("MODEL_PATH", "google/embeddinggemma-2") MAX_UPLOAD_MB = int(os.getenv("MAX_UPLOAD_MB", "20")) MAX_UPLOAD_BYTES = MAX_UPLOAD_MB * 1024 * 1024 MAX_BATCH = int(os.getenv("MAX_BATCH", "32")) BATCH_CHUNK = max(1, int(os.getenv("BATCH_CHUNK", "8"))) MAX_TOKENS = max(1, min(8192, int(os.getenv("MAX_TOKENS", "2048")))) MAX_TEXT_LENGTH = int(os.getenv("MAX_TEXT_LENGTH", "12000")) MAX_AUDIO_SECONDS = int(os.getenv("MAX_AUDIO_SECONDS", "60")) ALLOWED_DIMS = (128, 256, 512, 768) # Google-recommended task instructions. These are NOT just "SearchQuery: ...". TASK_PREFIXES = { "SearchQuery": "task: search result | query: {text}", "Document": "title: {title} | text: {text}", "QuestionAnswering": "task: question answering | query: {text}", "FactChecking": "task: fact checking | query: {text}", "CodeRetrieval": "task: code retrieval | query: {text}", "Classification": "task: classification | query: {text}", "Clustering": "task: clustering | query: {text}", "SentenceSimilarity": "task: sentence similarity | query: {text}", } _device = "cuda" if torch.cuda.is_available() else "cpu" # Avoid fp16: EmbeddingGemma 2's activations may be unstable in float16. _dtype = torch.bfloat16 if _device == "cuda" and torch.cuda.is_bf16_supported() else torch.float32 _processor = None _model = None _infer_lock = threading.Lock() @asynccontextmanager async def lifespan(app: FastAPI): global _processor, _model _processor = AutoProcessor.from_pretrained(MODEL_ID) _model = AutoModel.from_pretrained(MODEL_ID, dtype=_dtype).to(_device) _model.eval() yield _model = None _processor = None app = FastAPI( title="EmbeddingGemma 2 Multimodal API", version="2.0.0", description="API-only service for text, image, audio, and video embeddings.", lifespan=lifespan, docs_url=None, redoc_url=None, ) class ORJSONResponse(Response): media_type = "application/json" def render(self, content) -> bytes: return orjson.dumps(content) class EmbedRequest(BaseModel): text: str | None = Field(default=None, description="Text to embed (one input)") texts: list[str] | None = Field(default=None, description="Batch text inputs (max MAX_BATCH)") image_b64: str | None = Field(default=None, description="Base64 image, optionally a data URI") audio_b64: str | None = Field(default=None, description="Base64 WAV/FLAC/OGG audio") task_prefix: str | None = Field(default="SearchQuery", description="SearchQuery, Document, SentenceSimilarity, etc; or None") title: str | None = Field(default=None, description="Document title when task_prefix=Document") dimensions: Literal[128, 256, 512, 768] = 256 @model_validator(mode="after") def validate_request(self): if self.texts is not None: if self.text is not None or self.image_b64 is not None or self.audio_b64 is not None: raise ValueError("texts cannot be combined with text, image_b64, or audio_b64") if not (1 <= len(self.texts) <= MAX_BATCH): raise ValueError(f"texts must contain 1 to {MAX_BATCH} items") if any(not t.strip() or len(t) > MAX_TEXT_LENGTH for t in self.texts): raise ValueError("each text must be nonempty and within the text length limit") elif not any((self.text and self.text.strip(), self.image_b64, self.audio_b64)): raise ValueError("provide text, texts, image_b64, or audio_b64") if self.text is not None and len(self.text) > MAX_TEXT_LENGTH: raise ValueError("text is too long") if self.title is not None and len(self.title) > 500: raise ValueError("title is too long") if self.task_prefix not in (*TASK_PREFIXES, None, "None", "Raw"): raise ValueError("unsupported task_prefix") return self def format_text(text: str, task_prefix: str | None, title: str | None = None) -> str: if task_prefix in (None, "None", "Raw"): return text return TASK_PREFIXES[task_prefix].format(text=text, title=title or "none") def decode_base64(value: str, label: str) -> bytes: if value.startswith("data:"): if "," not in value: raise HTTPException(400, f"Invalid {label} data URI") value = value.split(",", 1)[1] # Base64 grows ~4/3; reject oversized uploads before decoding. if len(value) > ((MAX_UPLOAD_BYTES + 2) // 3) * 4 + 16: raise HTTPException(413, f"{label} is larger than {MAX_UPLOAD_MB} MB") try: decoded = base64.b64decode(value, validate=True) except (ValueError, binascii.Error) as exc: raise HTTPException(400, f"Invalid {label} base64") from exc if len(decoded) > MAX_UPLOAD_BYTES: raise HTTPException(413, f"{label} is larger than {MAX_UPLOAD_MB} MB") return decoded def load_image(contents: bytes) -> Image.Image: try: with Image.open(io.BytesIO(contents)) as opened: if opened.width * opened.height > 20_000_000: raise HTTPException(413, "Image exceeds 20 megapixels") opened.load() return opened.convert("RGB") except (UnidentifiedImageError, OSError, ValueError) as exc: raise HTTPException(400, "Cannot decode image") from exc def load_audio(contents: bytes) -> np.ndarray: try: with sf.SoundFile(io.BytesIO(contents)) as audio_file: sample_rate = int(audio_file.samplerate) if sample_rate <= 0 or len(audio_file) / sample_rate > MAX_AUDIO_SECONDS: raise HTTPException(413, f"Audio must be at most {MAX_AUDIO_SECONDS} seconds") audio = audio_file.read(dtype="float32", always_2d=True).mean(axis=1) except (sf.LibsndfileError, RuntimeError, ValueError) as exc: raise HTTPException(400, "Cannot decode audio; try WAV, FLAC, or OGG") from exc if len(audio) == 0: raise HTTPException(400, "Empty audio") if sample_rate != 16000: # Linear interpolation avoids a heavy resampling dependency. target_n = max(1, round(len(audio) * 16000 / sample_rate)) original_x = np.arange(len(audio), dtype=np.float64) target_x = np.arange(target_n, dtype=np.float64) * (sample_rate / 16000) audio = np.interp(target_x, original_x, audio).astype(np.float32) return audio def normalize_embeddings(output, attention_mask: torch.Tensor | None, dimensions: int) -> list[list[float]]: """Mask-aware mean pool, then truncate and L2 normalize in fp32.""" hidden = output.last_hidden_state.float() # (batch, tokens, 768) if attention_mask is None: pooled = hidden.mean(dim=1) else: mask = attention_mask.to(hidden.device).unsqueeze(-1).float() pooled = (hidden * mask).sum(dim=1) / mask.sum(dim=1).clamp(min=1.0) vec = pooled[:, :dimensions] vec = torch.nn.functional.normalize(vec, dim=1, eps=1e-12) return vec.cpu().tolist() def forward(inputs: dict, dimensions: int) -> list[list[float]]: if _model is None or _processor is None: raise HTTPException(503, "Model is still loading") with _infer_lock, torch.inference_mode(): model_inputs = {key: value.to(_device) if isinstance(value, torch.Tensor) else value for key, value in inputs.items()} out = _model(**model_inputs) return normalize_embeddings(out, model_inputs.get("attention_mask"), dimensions) def embed_texts(texts: list[str], task_prefix: str | None, title: str | None, dimensions: int): assert _processor is not None vectors = [] for offset in range(0, len(texts), BATCH_CHUNK): strings = [format_text(t, task_prefix, title) for t in texts[offset:offset + BATCH_CHUNK]] inputs = _processor(text=strings, return_tensors="pt", padding=True, truncation=True, max_length=MAX_TOKENS) vectors.extend(forward(inputs, dimensions)) return vectors def embed_single(dimensions: int, text: str | None = None, image: Image.Image | None = None, audio: np.ndarray | None = None, video_path: str | None = None, task_prefix: str | None = "SearchQuery", title: str | None = None) -> list[float]: assert _processor is not None kwargs = {"return_tensors": "pt"} has_media = image is not None or audio is not None or video_path is not None if text is not None and text.strip(): if has_media: # Direct multimodal processor calls require explicit placeholders # whenever text is supplied alongside media. tokens = [] if image is not None and "<|image|>" not in text: tokens.append("<|image|>") if audio is not None and "<|audio|>" not in text: tokens.append("<|audio|>") if video_path is not None and "<|video|>" not in text: tokens.append("<|video|>") kwargs["text"] = [" ".join(tokens + [text])] else: kwargs["text"] = [format_text(text, task_prefix, title)] if image is not None: kwargs["images"] = [image] if audio is not None: kwargs["audio"] = [audio] if video_path is not None: kwargs["videos"] = [video_path] kwargs["fps"] = 1 kwargs["max_frames"] = 8 inputs = _processor(**kwargs) return forward(inputs, dimensions)[0] @app.get("/", response_class=ORJSONResponse, tags=["Meta"]) def home(): return ORJSONResponse({ "status": "ready" if _model is not None else "loading", "engine": "multimodal-embeddinggemma-2", "model": MODEL_ID, "version": "2.0.0", "endpoints": { "health": "/health", "embed": "/embed", "embed_file": "/embed/file", "documentation": "/docs.MD", "openapi": "/openapi.json", }, }) @app.get("/health", response_class=ORJSONResponse, tags=["Meta"]) def health(): return ORJSONResponse({ "status": "ready" if _model is not None else "loading", "engine": "multimodal-embeddinggemma-2", "model": MODEL_ID, "device": _device, "dtype": str(_dtype).replace("torch.", ""), "dimensions": list(ALLOWED_DIMS), "max_batch": MAX_BATCH, "max_tokens": MAX_TOKENS, }) @app.get("/docs.MD", include_in_schema=False) def markdown_docs(): # Self-contained, copyable Markdown; no HTML frontend or docs.MD file needed. documentation = """# EmbeddingGemma 2 Multimodal API Base URL: `https://misukisu-embeddinggemma-2-api.hf.space` - `GET /` — JSON API metadata - `GET /health` — health and model readiness - `POST /embed` — one embedding or a batch of text embeddings - `POST /embed/file` — image, audio or video file (multipart/form-data) - `GET /docs.MD` — this Markdown reference - `GET /openapi.json` — machine-readable OpenAPI schema ## Text embedding ```bash curl -X POST https://misukisu-embeddinggemma-2-api.hf.space/embed \ -H 'Content-Type: application/json' \ -d '{"text":"Finland","task_prefix":"SearchQuery","dimensions":256}' ``` Single response: `{"embedding":[...],"dimensions":256}`. ## Batch embedding Use `{"texts":["Finland","Elon Musk"],"task_prefix":"Document","dimensions":256}`. Response: `{"embeddings":[[...],[...]],"dimensions":256,"count":2}`. ## Image and audio as JSON Send `image_b64` or `audio_b64` with base64-encoded media to `/embed`. Optional `text` can accompany the media. ## Image, video, or audio file ```bash curl -X POST https://misukisu-embeddinggemma-2-api.hf.space/embed/file \ -F 'file=@example.jpg' -F 'dimensions=256' ``` File response: `{"filename":"example.jpg","type":"image/jpeg","embedding":[...],"dimensions":256}`. Valid dimensions: 128, 256, 512, 768. Task prefixes: SearchQuery, Document, QuestionAnswering, FactChecking, CodeRetrieval, Classification, Clustering, SentenceSimilarity, Raw (or null). Task prefixes apply to text-only requests. """ return PlainTextResponse(documentation, media_type="text/markdown") @app.post("/embed", tags=["Embeddings"]) def embed_json(payload: EmbedRequest): """Accept one multimodal input or a batch of text inputs.""" if _model is None: raise HTTPException(503, "Model not ready") if payload.texts is not None: vectors = embed_texts(payload.texts, payload.task_prefix, payload.title, payload.dimensions) return ORJSONResponse({"embeddings": vectors, "dimensions": payload.dimensions, "count": len(vectors)}) image = load_image(decode_base64(payload.image_b64, "image")) if payload.image_b64 else None audio = load_audio(decode_base64(payload.audio_b64, "audio")) if payload.audio_b64 else None vector = embed_single(payload.dimensions, text=payload.text, image=image, audio=audio, task_prefix=payload.task_prefix, title=payload.title) return ORJSONResponse({"embedding": vector, "dimensions": payload.dimensions}) @app.post("/embed/file", tags=["Embeddings"]) async def embed_file( file: Annotated[UploadFile, File(description="An image, audio, or video file")], text: Annotated[str | None, Form()] = None, task_prefix: Annotated[str | None, Form()] = "SearchQuery", dimensions: Annotated[int, Form()] = 256, ): """Accept multipart/form-data media uploads; embed the file with optional text.""" if dimensions not in ALLOWED_DIMS: raise HTTPException(422, "dimensions must be 128, 256, 512, or 768") if task_prefix not in (*TASK_PREFIXES, None, "None", "Raw"): raise HTTPException(422, "unsupported task_prefix") if text is not None and len(text) > MAX_TEXT_LENGTH: raise HTTPException(413, "text too long") if _model is None: raise HTTPException(503, "Model not ready") contents = await file.read(MAX_UPLOAD_BYTES + 1) await file.close() if len(contents) > MAX_UPLOAD_BYTES: raise HTTPException(413, f"File exceeds {MAX_UPLOAD_MB} MB") if not contents: raise HTTPException(400, "Empty file") media_type = (file.content_type or "").lower().split(";", 1)[0] if not media_type.startswith(("image/", "audio/", "video/")): raise HTTPException(400, f"Unsupported media type: {media_type or 'unknown'}") if media_type.startswith("image/"): image = load_image(contents) vector = await run_in_threadpool(embed_single, dimensions, text=text, image=image, task_prefix=task_prefix) elif media_type.startswith("audio/"): audio = load_audio(contents) vector = await run_in_threadpool(embed_single, dimensions, text=text, audio=audio, task_prefix=task_prefix) else: # Let the model's native video processor handle video framing (not images=frames). extension = ".webm" if "webm" in media_type else ".mov" if "quicktime" in media_type else ".mp4" path = None try: with tempfile.NamedTemporaryFile(suffix=extension, delete=False) as tmp: tmp.write(contents) path = tmp.name try: vector = await run_in_threadpool(embed_single, dimensions, text=text, video_path=path, task_prefix=task_prefix) except (ValueError, OSError) as exc: raise HTTPException(400, "Cannot decode video") from exc finally: if path: Path(path).unlink(missing_ok=True) return ORJSONResponse({"filename": file.filename, "type": media_type, "embedding": vector, "dimensions": dimensions})