aakarshanr/apt_mia_processing
0
1from fastapi import APIRouter, UploadFile, File2import numpy as np3import cv24import io5from PIL import Image6 7from core.config import (8 IDX_TO_CLASS,9 THRESHOLD,10 LAST_CONV_LAYER,11 IG_STEPS12)13 14from preprocessing.image_preprocess import preprocess_image15from models.model_loader import load_model16 17from xai.gradcam_pp import gradcam18from xai.integrated_gradients import integrated_gradients, ig_to_heatmap19 20from utils.image_utils import image_array_to_base6421from utils.response_utils import success_response22 23 24# =========================25# ROUTER26# =========================27router = APIRouter()28 29# =========================30# LOAD MODEL ONCE31# =========================32model = load_model()33 34 35# =========================36# HEALTH CHECK37# =========================38@router.get("/")39def health():40 return {"status": "ok", "message": "API running 🔥"}41 42 43# =========================44# PREDICTION ENDPOINT45# =========================46@router.post("/predict")47async def predict(file: UploadFile = File(...)):48 try:49 # -----------------------------------------50 # READ IMAGE51 # -----------------------------------------52 image_bytes = await file.read()53 54 original_pil = Image.open(io.BytesIO(image_bytes)).convert("RGB")55 original_rgb = np.array(original_pil)56 original_bgr = cv2.cvtColor(original_rgb, cv2.COLOR_RGB2BGR)57 58 # -----------------------------------------59 # PREPROCESS60 # returns: {"input_layer_19": tensor}61 # -----------------------------------------62 inputs = preprocess_image(image_bytes)63 64 # Extract tensor for IG65 img_array = list(inputs.values())[0] # (1, H, W, 3)66 67 # -----------------------------------------68 # MODEL PREDICTION (SIGMOID)69 # -----------------------------------------70 pred = model.predict(inputs, verbose=0)71 sigmoid_value = float(np.squeeze(pred))72 73 if sigmoid_value < THRESHOLD:74 class_idx = 075 class_label = IDX_TO_CLASS[0]76 confidence = 1.0 - sigmoid_value77 else:78 class_idx = 179 class_label = IDX_TO_CLASS[1]80 confidence = sigmoid_value81 82 # -----------------------------------------83 # GRAD-CAM (BINARY SAFE)84 # -----------------------------------------85 cam = gradcam(86 model=model,87 inputs=inputs,88 last_conv_layer_name=LAST_CONV_LAYER89 )90 91 cam = cv2.resize(cam, (original_rgb.shape[1], original_rgb.shape[0]))92 cam = np.uint8(255 * cam)93 cam_color = cv2.applyColorMap(cam, cv2.COLORMAP_JET)94 95 gradcam_base64 = image_array_to_base64(cam_color)96 97 # -----------------------------------------98 # INTEGRATED GRADIENTS99 # -----------------------------------------100 ig_attr = integrated_gradients(101 model=model,102 img_array=img_array,103 class_idx=class_idx,104 steps=IG_STEPS105 )106 107 ig_img = ig_to_heatmap(ig_attr, original_bgr)108 ig_base64 = image_array_to_base64(ig_img)109 110 # -----------------------------------------111 # RESPONSE112 # -----------------------------------------113 return success_response(114 prediction=class_label,115 confidence=confidence,116 gradcam=gradcam_base64,117 integrated_gradients=ig_base64118 )119 120 except Exception as e:121 import traceback122 traceback.print_exc()123 return {124 "error": "Prediction failed",125 "details": str(e)126 }127 