Spaces:
Running
Running
Download app.py from misukisu/embeddinggemma-2-api: direct link, hf CLI and curl.
- Browser
- Download file 16.7 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
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() | |
| 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 | |
| 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] | |
| 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", | |
| }, | |
| }) | |
| 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, | |
| }) | |
| 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") | |
| 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}) | |
| 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}) | |