Spaces:
Sleeping
Sleeping
| """BioCLIP 2 zero-shot tree identification service. | |
| POST /identify — multipart image + JSON labels → top-5 softmax over candidates. | |
| Optional bearer auth via BIOCLIP_SERVICE_TOKEN (or BIOCLIP_HF_TOKEN). | |
| """ | |
| from __future__ import annotations | |
| import asyncio | |
| import hashlib | |
| import io | |
| import json | |
| import os | |
| import secrets | |
| from contextlib import asynccontextmanager | |
| from typing import Any | |
| import open_clip | |
| import torch | |
| import torch.nn.functional as F | |
| from fastapi import Depends, FastAPI, File, Form, Header, HTTPException, UploadFile | |
| from PIL import Image | |
| MODEL_ID = os.environ.get("BIOCLIP_MODEL_ID", "hf-hub:imageomics/bioclip-2").strip() | |
| _model: Any = None | |
| _preprocess: Any = None | |
| _tokenizer: Any = None | |
| _device: torch.device | None = None | |
| _infer_lock = asyncio.Lock() | |
| # Text encode dominates CPU cost for large closed sets — cache by label fingerprint. | |
| _text_cache_key: str | None = None | |
| _text_features: torch.Tensor | None = None | |
| def _expected_token() -> str | None: | |
| token = ( | |
| os.environ.get("BIOCLIP_SERVICE_TOKEN", "").strip() | |
| or os.environ.get("BIOCLIP_HF_TOKEN", "").strip() | |
| ) | |
| return token or None | |
| def require_bearer(authorization: str | None = Header(default=None)) -> None: | |
| expected = _expected_token() | |
| if not expected: | |
| return | |
| if not authorization or not authorization.startswith("Bearer "): | |
| raise HTTPException(status_code=401, detail="Missing Bearer token") | |
| got = authorization[len("Bearer ") :].strip() | |
| if not secrets.compare_digest(got, expected): | |
| raise HTTPException(status_code=401, detail="Invalid Bearer token") | |
| def _load_model() -> None: | |
| global _model, _preprocess, _tokenizer, _device | |
| _device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| model, _, preprocess = open_clip.create_model_and_transforms(MODEL_ID) | |
| model = model.to(_device) | |
| model.eval() | |
| _model = model | |
| _preprocess = preprocess | |
| _tokenizer = open_clip.get_tokenizer(MODEL_ID) | |
| async def lifespan(_app: FastAPI): | |
| _load_model() | |
| yield | |
| app = FastAPI(title="BioCLIP 2 Tree ID", lifespan=lifespan) | |
| def health() -> dict[str, Any]: | |
| return { | |
| "status": "ok", | |
| "model": MODEL_ID, | |
| "device": str(_device) if _device is not None else "loading", | |
| "auth_required": _expected_token() is not None, | |
| "text_cache": _text_cache_key is not None, | |
| } | |
| def _labels_fingerprint(label_list: list[str]) -> str: | |
| joined = "\n".join(name.strip() for name in label_list) | |
| return hashlib.sha256(joined.encode("utf-8")).hexdigest() | |
| def _text_features_for(label_list: list[str]) -> torch.Tensor: | |
| """Encode candidate labels once per distinct set (reuse across photos).""" | |
| global _text_cache_key, _text_features | |
| assert _model is not None and _tokenizer is not None and _device is not None | |
| key = _labels_fingerprint(label_list) | |
| if _text_cache_key == key and _text_features is not None: | |
| return _text_features | |
| texts = [f"a photo of {name.strip()}" for name in label_list] | |
| text_tokens = _tokenizer(texts).to(_device) | |
| with torch.no_grad(): | |
| feats = _model.encode_text(text_tokens) | |
| feats = F.normalize(feats, dim=-1) | |
| _text_cache_key = key | |
| _text_features = feats | |
| return feats | |
| def _identify_sync(raw: bytes, labels_json: str) -> dict[str, Any]: | |
| if _model is None or _preprocess is None or _tokenizer is None or _device is None: | |
| raise HTTPException(status_code=503, detail="Model not loaded") | |
| try: | |
| label_list = json.loads(labels_json) | |
| except json.JSONDecodeError as exc: | |
| raise HTTPException(status_code=400, detail="labels must be a JSON array") from exc | |
| if not isinstance(label_list, list) or not label_list: | |
| raise HTTPException(status_code=400, detail="labels must be a non-empty JSON array") | |
| if not all(isinstance(x, str) and x.strip() for x in label_list): | |
| raise HTTPException(status_code=400, detail="each label must be a non-empty string") | |
| # Soft upper bound to avoid OOM on CPU Spaces with huge closed sets. | |
| max_labels = int(os.environ.get("BIOCLIP_MAX_LABELS", "1000")) | |
| if len(label_list) > max_labels: | |
| raise HTTPException( | |
| status_code=400, | |
| detail=f"at most {max_labels} labels allowed (got {len(label_list)})", | |
| ) | |
| if not raw: | |
| raise HTTPException(status_code=400, detail="empty image") | |
| try: | |
| pil = Image.open(io.BytesIO(raw)).convert("RGB") | |
| except Exception as exc: | |
| raise HTTPException(status_code=400, detail=f"invalid image: {exc}") from exc | |
| image_tensor = _preprocess(pil).unsqueeze(0).to(_device) | |
| text_features = _text_features_for(label_list) | |
| with torch.no_grad(): | |
| image_features = _model.encode_image(image_tensor) | |
| image_features = F.normalize(image_features, dim=-1) | |
| logits = (image_features @ text_features.T).squeeze(0) | |
| probs = F.softmax(logits.float(), dim=-1) | |
| ranked = sorted( | |
| ( | |
| {"latin": label_list[i].strip(), "probability": float(probs[i].item())} | |
| for i in range(len(label_list)) | |
| ), | |
| key=lambda x: x["probability"], | |
| reverse=True, | |
| ) | |
| # Absolute closed-set softmax — often tiny across ~800 taxa; do not renorm | |
| # (renorm / genus rollups invent false confidence when the distribution is flat). | |
| return {"top5": ranked[:5]} | |
| async def identify( | |
| image: UploadFile = File(...), | |
| labels: str = Form(...), | |
| ) -> dict[str, Any]: | |
| """Zero-shot classify an image against candidate latin names only.""" | |
| raw = await image.read() | |
| # One inference at a time — CPU Spaces OOM/timeout when Vercel fans out 6 photos. | |
| async with _infer_lock: | |
| return await asyncio.to_thread(_identify_sync, raw, labels) | |
| def main() -> None: | |
| import uvicorn | |
| # Hugging Face Spaces inject SPACE_ID and expect port 7860. | |
| default_port = "7860" if os.environ.get("SPACE_ID") else "8090" | |
| port = int(os.environ.get("PORT", default_port)) | |
| uvicorn.run("main:app", host="0.0.0.0", port=port, reload=False) | |
| if __name__ == "__main__": | |
| main() | |