import io import os import time import logging import torch import numpy as np from PIL import Image from fastapi import FastAPI, File, UploadFile, HTTPException from fastapi.responses import JSONResponse from wildlife_tools.features import DeepFeatures from timm import create_model from torchvision import transforms from torchvision.models import mobilenet_v2, MobileNet_V2_Weights DOG_CLASS_MIN = 151 DOG_CLASS_MAX = 268 DOG_CLASSIFIER_TOP_K = int(os.environ.get("DOG_CLASSIFIER_TOP_K", "3")) logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) app = FastAPI(title="Jeeva Embedding Service") # Global state: model loaded once at startup _extractor = None _transform = None _dog_classifier = None _dog_transform = None _last_embed_ts = None @app.on_event("startup") async def load_model(): global _extractor, _transform, _dog_classifier, _dog_transform logger.info("Loading MobileNetV2 for dog classification...") weights = MobileNet_V2_Weights.IMAGENET1K_V1 _dog_classifier = mobilenet_v2(weights=weights) _dog_classifier.eval() _dog_transform = weights.transforms() logger.info("Loading MegaDescriptor-T-224 from HF Hub...") t0 = time.time() backbone = create_model( "hf-hub:BVRA/MegaDescriptor-T-224", pretrained=True, num_classes=0, # remove classification head, output features only ) backbone.eval() _extractor = DeepFeatures(backbone, device="cpu", batch_size=1) _transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) logger.info(f"Model loaded in {time.time() - t0:.1f}s") def _crop_face(img: Image.Image) -> Image.Image: """Top 40% of image height — heuristic face region for ground-level dog photos.""" w, h = img.size return img.crop((0, 0, w, int(h * 0.4))) def _embed(img: Image.Image) -> list[float]: """Run one image through the extractor and return a flat float list.""" tensor = _transform(img.convert("RGB")).unsqueeze(0) # [1, 3, 224, 224] with torch.no_grad(): feat = backbone_forward(tensor) return feat.squeeze(0).tolist() def backbone_forward(tensor: torch.Tensor) -> torch.Tensor: """Forward pass through the backbone (DeepFeatures wraps batch lists, so we do it directly).""" global _extractor return _extractor.model(tensor) @app.post("/embed") async def embed(file: UploadFile = File(...)): global _extractor, _last_embed_ts, _dog_classifier, _dog_transform if _extractor is None or _dog_classifier is None: raise HTTPException(status_code=503, detail="Model not loaded yet") contents = await file.read() try: img = Image.open(io.BytesIO(contents)).convert("RGB") except Exception: raise HTTPException(status_code=400, detail="Could not decode image") dog_tensor = _dog_transform(img).unsqueeze(0) with torch.no_grad(): preds = _dog_classifier(dog_tensor) top_k_indices = torch.topk(preds, DOG_CLASSIFIER_TOP_K, dim=1).indices.squeeze(0).tolist() is_dog = any(DOG_CLASS_MIN <= idx <= DOG_CLASS_MAX for idx in top_k_indices) if not is_dog: return JSONResponse({ "is_dog": False }) face_crop = _crop_face(img) body_crop = img # full image as body region face_vec = _embed(face_crop) body_vec = _embed(body_crop) _last_embed_ts = time.time() return JSONResponse({ "is_dog": True, "face_embedding": face_vec, "body_embedding": body_vec, "dims": len(face_vec), }) @app.get("/health") async def health(): return { "status": "ok", "model_loaded": _extractor is not None, "last_embed_ts": _last_embed_ts, }