justanitin's picture
Upload 2 files
e7f7cda verified
Raw
History Blame Contribute Delete
5.31 kB
import os
import io
import json
import base64
import asyncio
import numpy as np
from PIL import Image
from contextlib import asynccontextmanager
from fastapi import FastAPI, HTTPException, status, Header, Depends
from fastapi.responses import JSONResponse
from pydantic import BaseModel
import onnxruntime as ort
# Add parent dir to path so we can import registry module
import sys
from pathlib import Path
sys.path.append(str(Path(__file__).parent.parent))
from registry.download_model import download_artifacts, TARGET_DIR, VERSION
# Global state
is_model_ready = False
ort_session = None
class_labels = []
temperature = 1.0
INTERNAL_SECRET = os.environ.get("INTERNAL_SECRET")
if not INTERNAL_SECRET:
raise ValueError("CRITICAL: INTERNAL_SECRET environment variable is missing.")
@asynccontextmanager
async def lifespan(app: FastAPI):
global is_model_ready, ort_session, class_labels, temperature
# 1. Download artifacts
# NOTE: In production on HF Spaces, HF_TOKEN must be set if repo is private
# We run this in a thread to not block the event loop
print("Initializing Hugging Face Space Lifespan...")
try:
await asyncio.to_thread(download_artifacts)
except Exception as e:
print(f"Failed to download artifacts: {e}")
# 2. Load the ONNX model
onnx_path = TARGET_DIR / f"cropguard_{VERSION}.onnx"
if not onnx_path.exists():
print(f"Warning: Model not found at {onnx_path}. Space will fail health checks.")
else:
# Load ONNX model
providers = ['CPUExecutionProvider']
ort_session = ort.InferenceSession(str(onnx_path), providers=providers)
# Load labels
json_path = TARGET_DIR / f"cropguard_{VERSION}.json"
if json_path.exists():
with open(json_path, 'r') as f:
metadata = json.load(f)
class_labels = metadata.get("classes", [])
# Load temperature (optional)
temp_path = TARGET_DIR / "temperature.json"
if temp_path.exists():
with open(temp_path, 'r') as f:
temperature = json.load(f).get("temperature", 1.0)
is_model_ready = True
print("Model loaded successfully. Engine ready.")
yield
app = FastAPI(lifespan=lifespan)
class InferenceRequest(BaseModel):
image_b64: str
def verify_internal_secret(x_internal_secret: str = Header(None)):
if x_internal_secret != INTERNAL_SECRET:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Forbidden: Invalid or missing X-Internal-Secret header"
)
return x_internal_secret
@app.get("/health")
async def health_check():
if not is_model_ready:
return JSONResponse(status_code=503, content={"status": "loading"})
return {"status": "ready"}
def preprocess_image(image_b64: str):
# Decode base64
if "," in image_b64:
image_b64 = image_b64.split(",")[1]
image_bytes = base64.b64decode(image_b64)
image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
# Resize to 224x224 (ConvNeXt standard)
image = image.resize((224, 224), Image.Resampling.BILINEAR)
# Convert to numpy array and normalize
# ConvNeXt ImageNet Mean: [0.485, 0.456, 0.406]
# ConvNeXt ImageNet Std: [0.229, 0.224, 0.225]
img_arr = np.array(image, dtype=np.float32) / 255.0
mean = np.array([0.485, 0.456, 0.406], dtype=np.float32)
std = np.array([0.229, 0.224, 0.225], dtype=np.float32)
img_arr = (img_arr - mean) / std
# HWC to CHW format
img_arr = np.transpose(img_arr, (2, 0, 1))
# Add batch dimension
img_arr = np.expand_dims(img_arr, axis=0)
return img_arr
def run_inference(img_arr: np.ndarray):
input_name = ort_session.get_inputs()[0].name
outputs = ort_session.run(None, {input_name: img_arr})
logits = outputs[0][0]
# Apply temperature scaling
scaled_logits = logits / temperature
# Softmax
exp_logits = np.exp(scaled_logits - np.max(scaled_logits))
probs = exp_logits / exp_logits.sum()
pred_idx = int(np.argmax(probs))
confidence = float(probs[pred_idx])
label = class_labels[pred_idx] if pred_idx < len(class_labels) else f"class_{pred_idx}"
return label, confidence
@app.post("/infer", dependencies=[Depends(verify_internal_secret)])
async def infer(payload: InferenceRequest):
if not is_model_ready:
return JSONResponse(status_code=503, content={"status": "loading"})
try:
# 1. Preprocess
img_arr = await asyncio.to_thread(preprocess_image, payload.image_b64)
# 2. Run Inference
label, confidence = await asyncio.to_thread(run_inference, img_arr)
# 3. Generate Heatmap (Mock for now, as True CAM requires PyTorch hooks)
mock_heatmap = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8z8BQDwAEhQGAhKmMIQAAAABJRU5ErkJggg=="
return {
"raw_label": label,
"confidence": confidence,
"heatmap_b64": mock_heatmap
}
except Exception as e:
print(f"Inference error: {e}")
raise HTTPException(status_code=500, detail="Inference failed")