jeeva-embed / app.py
shankzakaneo's picture
feat: add MobileNetV2 dog-validation gate before embedding
0978eb3
Raw
History Blame Contribute Delete
3.85 kB
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,
}