Spaces:
Paused
Paused
| import os | |
| # ZeroGPU and library caches must be configured before importing spaces/torch. | |
| os.environ.setdefault("HF_HOME", os.path.expanduser("~/.cache/huggingface")) | |
| os.environ.setdefault("HF_MODULES_CACHE", "/tmp/hf_modules") | |
| os.environ.setdefault("MPLCONFIGDIR", "/tmp/matplotlib") | |
| os.environ.setdefault("GRADIO_ANALYTICS_ENABLED", "False") | |
| os.environ.setdefault("GRADIO_SSR_MODE", "false") | |
| os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") | |
| import spaces | |
| import logging | |
| import threading | |
| import time | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| from typing import Any | |
| import gradio as gr | |
| import torch | |
| import torch.nn.functional as F | |
| from fastapi.responses import HTMLResponse | |
| from fastapi.staticfiles import StaticFiles | |
| from gradio.data_classes import FileData | |
| from sentence_transformers import SentenceTransformer | |
| MODEL_ID = "tencent/WeMM-Embedding-9B" | |
| ROOT = Path(__file__).resolve().parent | |
| ASSET_DIR = ROOT / "assets" | |
| FRONTEND_DIR = ROOT / "frontend" | |
| MATRYOSHKA_DIMS = (64, 128, 256, 512, 1024, 2048, 4096) | |
| MAX_CUSTOM_TEXTS = 6 | |
| MAX_CUSTOM_MEDIA = 6 | |
| logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(levelname)s | %(message)s") | |
| LOGGER = logging.getLogger("wemm-space") | |
| class Candidate: | |
| key: str | |
| title: str | |
| kind: str | |
| description: str | |
| payload: Any | |
| media_path: str | None = None | |
| SHOWCASE: tuple[Candidate, ...] = ( | |
| Candidate( | |
| "llama4", | |
| "Llama 4 model card", | |
| "visual document", | |
| "A dense model-card screenshot describing the Scout and Maverick variants.", | |
| str(ASSET_DIR / "llama4_hgf.png"), | |
| str(ASSET_DIR / "llama4_hgf.png"), | |
| ), | |
| Candidate( | |
| "qwen-omni", | |
| "Qwen2.5-Omni overview", | |
| "visual document", | |
| "A model page covering omni-modal perception, speech, and video capabilities.", | |
| str(ASSET_DIR / "qwen2.5omni_hgf.png"), | |
| str(ASSET_DIR / "qwen2.5omni_hgf.png"), | |
| ), | |
| Candidate( | |
| "likelihood-contour", | |
| "Scientific contour plot", | |
| "figure", | |
| "An orange likelihood contour plotted against gamma and log-scaled tau over mass.", | |
| str(ASSET_DIR / "doc1.jpg"), | |
| str(ASSET_DIR / "doc1.jpg"), | |
| ), | |
| Candidate( | |
| "budget-1971", | |
| "1971 budget infographic", | |
| "visual document", | |
| "A historical chart comparing US outlays, including natural resources spending.", | |
| str(ASSET_DIR / "doc2.jpg"), | |
| str(ASSET_DIR / "doc2.jpg"), | |
| ), | |
| Candidate( | |
| "scoring-rules", | |
| "Proper scoring rules paper", | |
| "visual document", | |
| "An academic page with lemmas, an algorithm, equations, and references.", | |
| str(ASSET_DIR / "doc3.jpg"), | |
| str(ASSET_DIR / "doc3.jpg"), | |
| ), | |
| Candidate( | |
| "road-safety", | |
| "Road-safety assessment", | |
| "visual document", | |
| "An environmental assessment page about driver training, signs, and road closures.", | |
| str(ASSET_DIR / "doc4.jpg"), | |
| str(ASSET_DIR / "doc4.jpg"), | |
| ), | |
| Candidate( | |
| "mapo-tofu", | |
| "Mapo tofu in motion", | |
| "video", | |
| "A short cooking clip showing the preparation of the Sichuan tofu dish.", | |
| str(ASSET_DIR / "mapo_tofu.mp4"), | |
| str(ASSET_DIR / "mapo_tofu.mp4"), | |
| ), | |
| Candidate( | |
| "zhajiang-noodles", | |
| "Zhajiang noodles in motion", | |
| "video", | |
| "A short cooking clip showing noodles with a savory fermented-bean sauce.", | |
| str(ASSET_DIR / "zhajiang_noodle.mp4"), | |
| str(ASSET_DIR / "zhajiang_noodle.mp4"), | |
| ), | |
| Candidate( | |
| "vector-search", | |
| "How vector search works", | |
| "text", | |
| "Dense retrieval maps queries and documents into one normalized vector space, then ranks candidates by cosine similarity.", | |
| "Dense retrieval maps queries and documents into one normalized vector space, then ranks candidates by cosine similarity.", | |
| ), | |
| Candidate( | |
| "night-train", | |
| "A quiet journey", | |
| "text", | |
| "夜行列车穿过雨中的城市,车窗映出霓虹灯和安静的乘客。", | |
| "夜行列车穿过雨中的城市,车窗映出霓虹灯和安静的乘客。", | |
| ), | |
| ) | |
| LOGGER.info("Loading %s on CPU", MODEL_ID) | |
| MODEL = SentenceTransformer( | |
| MODEL_ID, | |
| trust_remote_code=True, | |
| device="cpu", | |
| model_kwargs={"dtype": torch.bfloat16, "low_cpu_mem_usage": True}, | |
| ) | |
| MODEL.eval() | |
| torch.set_grad_enabled(False) | |
| LOGGER.info("Model loaded; waiting for a ZeroGPU allocation") | |
| _CACHE_LOCK = threading.Lock() | |
| _INFERENCE_LOCK = threading.Lock() | |
| _SHOWCASE_EMBEDDINGS: torch.Tensor | None = None | |
| def _file_path(value: Any) -> str | None: | |
| """Normalize Gradio file values across browser and programmatic clients.""" | |
| if value is None: | |
| return None | |
| if isinstance(value, (str, Path)): | |
| return str(value) | |
| if isinstance(value, dict): | |
| path = value.get("path") or value.get("name") | |
| return str(path) if path else None | |
| path = getattr(value, "path", None) or getattr(value, "name", None) | |
| return str(path) if path else None | |
| def _is_video(path: str) -> bool: | |
| return Path(path.split("?", 1)[0]).suffix.lower() in { | |
| ".mp4", ".webm", ".mov", ".mkv", ".avi", ".mpeg", ".mpg" | |
| } | |
| def _multimodal_payload(text: str | None, image: Any, video: Any) -> tuple[Any, str]: | |
| text = (text or "").strip() | |
| image_path = _file_path(image) | |
| video_path = _file_path(video) | |
| if image_path and video_path: | |
| raise gr.Error("Choose one visual query: an image or a video, not both.") | |
| if image_path: | |
| return ({"image": image_path, "text": text}, "image + text") if text else (image_path, "image") | |
| if video_path: | |
| return ({"video": video_path, "text": text}, "video + text") if text else (video_path, "video") | |
| if text: | |
| return text, "text" | |
| raise gr.Error("Add a text, image, or video query to begin.") | |
| def _parse_text_candidates(raw: str | None) -> list[Candidate]: | |
| candidates: list[Candidate] = [] | |
| lines = [line.strip() for line in (raw or "").splitlines() if line.strip()] | |
| if len(lines) > MAX_CUSTOM_TEXTS: | |
| raise gr.Error(f"Use at most {MAX_CUSTOM_TEXTS} custom text candidates.") | |
| for index, line in enumerate(lines, start=1): | |
| if "::" in line: | |
| title, body = (part.strip() for part in line.split("::", 1)) | |
| title = title or f"Custom text {index}" | |
| body = body or title | |
| else: | |
| title, body = f"Custom text {index}", line | |
| candidates.append(Candidate(f"custom-text-{index}", title[:80], "text", body[:240], body)) | |
| return candidates | |
| def _parse_media_candidates(raw: Any) -> list[Candidate]: | |
| candidates: list[Candidate] = [] | |
| items = raw or [] | |
| if len(items) > MAX_CUSTOM_MEDIA: | |
| raise gr.Error(f"Upload at most {MAX_CUSTOM_MEDIA} candidate media files.") | |
| for index, item in enumerate(items, start=1): | |
| path = _file_path(item) | |
| if not path: | |
| continue | |
| kind = "video" if _is_video(path) else "image" | |
| candidates.append( | |
| Candidate( | |
| f"custom-media-{index}", f"Uploaded {kind} {index}", kind, | |
| f"User-supplied {kind} candidate.", path, path, | |
| ) | |
| ) | |
| return candidates | |
| def _encode(items: list[Any], *, query: bool) -> torch.Tensor: | |
| method = MODEL.encode_query if query else MODEL.encode_document | |
| with torch.inference_mode(): | |
| embeddings = method( | |
| items, batch_size=1, convert_to_tensor=True, | |
| normalize_embeddings=True, show_progress_bar=False, | |
| ) | |
| if embeddings.ndim == 1: | |
| embeddings = embeddings.unsqueeze(0) | |
| return embeddings.float().cpu() | |
| def _truncate_normalize(embeddings: torch.Tensor, dimension: int) -> torch.Tensor: | |
| return F.normalize(embeddings[..., :dimension], p=2, dim=-1) | |
| def _dimension_scores(query: torch.Tensor, documents: torch.Tensor) -> dict[int, list[float]]: | |
| scores: dict[int, list[float]] = {} | |
| for dimension in MATRYOSHKA_DIMS: | |
| query_d = _truncate_normalize(query, dimension) | |
| docs_d = _truncate_normalize(documents, dimension) | |
| scores[dimension] = (query_d @ docs_d.T).squeeze(0).tolist() | |
| return scores | |
| def _fingerprint(vector: torch.Tensor, bar_count: int = 96) -> list[float]: | |
| values = vector.detach().float().flatten() | |
| chunks = torch.tensor_split(values, min(bar_count, values.numel())) | |
| return [round(float(chunk.mean()), 7) for chunk in chunks] | |
| def _public_media(candidate: Candidate) -> str | None: | |
| if not candidate.media_path: | |
| return None | |
| try: | |
| relative = Path(candidate.media_path).resolve().relative_to(ASSET_DIR.resolve()) | |
| except ValueError: | |
| return None | |
| return f"/assets/{relative.as_posix()}" | |
| def _search_duration(*args: Any, **kwargs: Any) -> int: | |
| """Budget more time for the cold corpus pass and video queries.""" | |
| query_payload = args[0] if args else None | |
| custom_candidates = args[2] if len(args) > 2 else [] | |
| include_showcase = bool(args[3]) if len(args) > 3 else True | |
| query_has_video = ( | |
| isinstance(query_payload, dict) and bool(query_payload.get("video")) | |
| ) or (isinstance(query_payload, str) and _is_video(query_payload)) | |
| has_uploaded_video = any(item.kind == "video" for item in custom_candidates or []) | |
| if include_showcase and _SHOWCASE_EMBEDDINGS is None: | |
| return 30 | |
| return 30 if query_has_video or has_uploaded_video else 15 | |
| def _run_search( | |
| query_payload: Any, | |
| query_kind: str, | |
| custom_candidates: list[Candidate], | |
| include_showcase: bool, | |
| dimension: int, | |
| ) -> tuple[torch.Tensor, torch.Tensor, list[Candidate], float, bool]: | |
| global _SHOWCASE_EMBEDDINGS | |
| started = time.perf_counter() | |
| was_cold = include_showcase and _SHOWCASE_EMBEDDINGS is None | |
| MODEL.to("cuda") | |
| try: | |
| query_embedding = _encode([query_payload], query=True) | |
| document_blocks: list[torch.Tensor] = [] | |
| candidates: list[Candidate] = [] | |
| if include_showcase: | |
| with _CACHE_LOCK: | |
| cached = _SHOWCASE_EMBEDDINGS | |
| if cached is None: | |
| encoded = _encode([item.payload for item in SHOWCASE], query=False) | |
| with _CACHE_LOCK: | |
| if _SHOWCASE_EMBEDDINGS is None: | |
| _SHOWCASE_EMBEDDINGS = encoded | |
| cached = _SHOWCASE_EMBEDDINGS | |
| document_blocks.append(cached) | |
| candidates.extend(SHOWCASE) | |
| if custom_candidates: | |
| custom_embeddings = _encode([item.payload for item in custom_candidates], query=False) | |
| document_blocks.append(custom_embeddings) | |
| candidates.extend(custom_candidates) | |
| if not document_blocks: | |
| raise gr.Error("Include the showcase universe or add at least one candidate.") | |
| documents = torch.cat(document_blocks, dim=0) | |
| return query_embedding, documents, candidates, time.perf_counter() - started, was_cold | |
| finally: | |
| MODEL.to("cpu") | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| def _pair_duration(*args: Any, **kwargs: Any) -> int: | |
| has_video = any( | |
| (isinstance(value, dict) and bool(value.get("video"))) | |
| or (isinstance(value, str) and _is_video(value)) | |
| for value in args[:2] | |
| ) | |
| return 30 if has_video else 15 | |
| def _run_pair( | |
| query_payload: Any, | |
| candidate_payload: Any, | |
| ) -> tuple[torch.Tensor, torch.Tensor, float]: | |
| started = time.perf_counter() | |
| MODEL.to("cuda") | |
| try: | |
| query = _encode([query_payload], query=True) | |
| candidate = _encode([candidate_payload], query=False) | |
| return query, candidate, time.perf_counter() - started | |
| finally: | |
| MODEL.to("cpu") | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| def _interpret_score(score: float) -> tuple[str, str]: | |
| if score >= 0.75: | |
| return "high alignment", "These inputs occupy a very similar region for this retrieval model." | |
| if score >= 0.50: | |
| return "meaningful alignment", "The model sees a substantial semantic relationship." | |
| if score >= 0.25: | |
| return "weak alignment", "There is some overlap, but stronger candidates may rank above it." | |
| return "low alignment", "The model places these inputs relatively far apart." | |
| app = gr.Server( | |
| title="WeMM Semantic Universe API", | |
| summary="Custom multimodal retrieval studio powered by tencent/WeMM-Embedding-9B", | |
| version="2.0.0", | |
| ) | |
| def search_api( | |
| query_text: str = "", | |
| query_image: FileData | None = None, | |
| query_video: FileData | None = None, | |
| custom_texts: str = "", | |
| candidate_media: list[FileData] | None = None, | |
| include_showcase: bool = True, | |
| dimension: int = 1024, | |
| ) -> dict[str, Any]: | |
| """Search a mixed media collection and return structured visualization data.""" | |
| dimension = int(dimension) | |
| if dimension not in MATRYOSHKA_DIMS: | |
| raise gr.Error("Choose one of the model's native Matryoshka dimensions.") | |
| query_payload, query_kind = _multimodal_payload(query_text, query_image, query_video) | |
| custom_candidates = _parse_text_candidates(custom_texts) + _parse_media_candidates(candidate_media) | |
| with _INFERENCE_LOCK: | |
| query, documents, candidates, elapsed, was_cold = _run_search( | |
| query_payload, query_kind, custom_candidates, bool(include_showcase), dimension, | |
| ) | |
| dimension_map = _dimension_scores(query, documents) | |
| chosen_scores = dimension_map[dimension] | |
| order = sorted(range(len(candidates)), key=lambda index: chosen_scores[index], reverse=True) | |
| ranked = [ | |
| { | |
| "rank": rank, | |
| "key": candidates[index].key, | |
| "title": candidates[index].title, | |
| "kind": candidates[index].kind, | |
| "description": candidates[index].description, | |
| "score": round(float(chosen_scores[index]), 7), | |
| "media_url": _public_media(candidates[index]), | |
| } | |
| for rank, index in enumerate(order, start=1) | |
| ] | |
| curve_series = [ | |
| { | |
| "key": candidates[index].key, | |
| "title": candidates[index].title, | |
| "values": [round(float(dimension_map[dim][index]), 7) for dim in MATRYOSHKA_DIMS], | |
| } | |
| for index in order[:4] | |
| ] | |
| return { | |
| "mode": "search", | |
| "model": MODEL_ID, | |
| "query_kind": query_kind, | |
| "dimension": dimension, | |
| "dimensions": list(MATRYOSHKA_DIMS), | |
| "candidate_count": len(candidates), | |
| "elapsed_seconds": round(elapsed, 3), | |
| "cache_state": "created" if was_cold else "reused" if include_showcase else "not requested", | |
| "top_match": ranked[0], | |
| "rankings": ranked, | |
| "dimension_series": curve_series, | |
| "query_fingerprint": _fingerprint(query[0, :dimension]), | |
| "full_embedding_dimension": int(query.shape[-1]), | |
| "query_vector_preview": [round(float(value), 7) for value in query[0, :12]], | |
| "l2_norm_after_truncation": round(float(_truncate_normalize(query, dimension).norm()), 7), | |
| "note": "Cosine similarity is a ranking signal, not a calibrated probability.", | |
| } | |
| def compare_api( | |
| query_text: str = "", | |
| query_image: FileData | None = None, | |
| query_video: FileData | None = None, | |
| candidate_text: str = "", | |
| candidate_image: FileData | None = None, | |
| candidate_video: FileData | None = None, | |
| dimension: int = 1024, | |
| ) -> dict[str, Any]: | |
| """Compare two supported inputs and return structured visualization data.""" | |
| dimension = int(dimension) | |
| if dimension not in MATRYOSHKA_DIMS: | |
| raise gr.Error("Choose one of the model's native Matryoshka dimensions.") | |
| query_payload, query_kind = _multimodal_payload(query_text, query_image, query_video) | |
| candidate_payload, candidate_kind = _multimodal_payload( | |
| candidate_text, candidate_image, candidate_video | |
| ) | |
| with _INFERENCE_LOCK: | |
| query, candidate, elapsed = _run_pair(query_payload, candidate_payload) | |
| scores = { | |
| dim: float((_truncate_normalize(query, dim) @ _truncate_normalize(candidate, dim).T).item()) | |
| for dim in MATRYOSHKA_DIMS | |
| } | |
| selected_score = scores[dimension] | |
| label, explanation = _interpret_score(selected_score) | |
| return { | |
| "mode": "compare", | |
| "model": MODEL_ID, | |
| "query_kind": query_kind, | |
| "candidate_kind": candidate_kind, | |
| "dimension": dimension, | |
| "dimensions": list(MATRYOSHKA_DIMS), | |
| "selected_score": round(selected_score, 7), | |
| "label": label, | |
| "explanation": explanation, | |
| "scores": [round(scores[dim], 7) for dim in MATRYOSHKA_DIMS], | |
| "query_fingerprint": _fingerprint(query[0, :dimension]), | |
| "candidate_fingerprint": _fingerprint(candidate[0, :dimension]), | |
| "elapsed_seconds": round(elapsed, 3), | |
| "note": "Interpret thresholds relative to a task-specific candidate set.", | |
| } | |
| async def homepage() -> HTMLResponse: | |
| return HTMLResponse((FRONTEND_DIR / "index.html").read_text(encoding="utf-8")) | |
| async def health() -> dict[str, str]: | |
| return {"status": "ok", "model": MODEL_ID, "interface": "gradio-server"} | |
| app.mount("/assets", StaticFiles(directory=str(ASSET_DIR)), name="assets") | |
| app.mount("/ui", StaticFiles(directory=str(FRONTEND_DIR)), name="ui") | |
| # Hugging Face Spaces expects the application object under this conventional name. | |
| demo = app | |
| if __name__ == "__main__": | |
| app.launch(show_error=True) | |