misukisu's picture
Update app.py
87e67a3 verified
Raw History Blame Contribute Delete
16.7 kB
"""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})