| """ |
| Pneumonia Detection API |
| FastAPI backend for dual-model pneumonia classification from chest X-rays |
| """ |
|
|
| from fastapi import FastAPI, File, UploadFile, HTTPException |
| from fastapi.middleware.cors import CORSMiddleware |
| from fastapi.responses import JSONResponse |
| import uvicorn |
| import numpy as np |
| import time |
| import io |
| from PIL import Image |
| import cv2 |
|
|
| |
| from utils.preprocess_svm import preprocess_for_svm, extract_svm_features |
| from utils.preprocess_cnn import preprocess_for_cnn |
| from utils.model_loader import load_models |
|
|
| |
| app = FastAPI( |
| title="Pneumonia Detection API", |
| description="Dual-model (SVM + DenseNet121) pneumonia classification API", |
| version="1.0.0" |
| ) |
|
|
| |
| app.add_middleware( |
| CORSMiddleware, |
| allow_origins=["*"], |
| allow_credentials=True, |
| allow_methods=["*"], |
| allow_headers=["*"], |
| ) |
|
|
| |
| MODELS = {} |
|
|
| @app.on_event("startup") |
| async def startup_event(): |
| """Load models on server startup""" |
| global MODELS |
| print("Loading models...") |
| MODELS = load_models() |
| print("✅ Models loaded successfully!") |
| print(f" - SVM model: {type(MODELS['svm']).__name__}") |
| print(f" - CNN model: {type(MODELS['cnn']).__name__}") |
| print(f" - Device: {MODELS['device']}") |
|
|
|
|
| @app.get("/") |
| async def root(): |
| """Health check endpoint""" |
| return { |
| "status": "healthy", |
| "message": "Pneumonia Detection API is running", |
| "models_loaded": len(MODELS) > 0, |
| "version": "1.0.0" |
| } |
|
|
|
|
| @app.get("/health") |
| async def health_check(): |
| """Detailed health check""" |
| return { |
| "status": "healthy", |
| "models": { |
| "svm": MODELS.get("svm") is not None, |
| "cnn": MODELS.get("cnn") is not None, |
| }, |
| "device": str(MODELS.get("device", "unknown")) |
| } |
|
|
|
|
| @app.post("/predict") |
| async def predict_pneumonia(file: UploadFile = File(...)): |
| """ |
| Predict pneumonia from chest X-ray image |
| |
| Args: |
| file: Uploaded image file (JPG, JPEG, PNG) |
| |
| Returns: |
| JSON with predictions from both models |
| """ |
| |
| |
| if not file.content_type.startswith("image/"): |
| raise HTTPException( |
| status_code=400, |
| detail="Invalid file type. Please upload an image (JPG, JPEG, PNG)" |
| ) |
| |
| try: |
| |
| contents = await file.read() |
| image = Image.open(io.BytesIO(contents)) |
| |
| |
| image_np = np.array(image.convert('RGB')) |
| |
| |
| image_bgr = cv2.cvtColor(image_np, cv2.COLOR_RGB2BGR) |
| |
| |
| |
| |
| start_svm = time.time() |
| |
| |
| svm_preprocessed = preprocess_for_svm(image_bgr) |
| |
| |
| svm_features = extract_svm_features(svm_preprocessed) |
| |
| |
| svm_features = svm_features.reshape(1, -1) |
| |
| |
| svm_prediction = MODELS['svm'].predict(svm_features)[0] |
| svm_proba = MODELS['svm'].predict_proba(svm_features)[0] |
| |
| svm_time = time.time() - start_svm |
| |
| |
| |
| |
| start_cnn = time.time() |
| |
| |
| cnn_input = preprocess_for_cnn(image_np) |
| cnn_input = cnn_input.unsqueeze(0).to(MODELS['device']) |
| |
| |
| MODELS['cnn'].eval() |
| import torch |
| with torch.no_grad(): |
| outputs = MODELS['cnn'](cnn_input) |
| probabilities = torch.softmax(outputs, dim=1) |
| cnn_prediction = outputs.argmax(dim=1).item() |
| cnn_confidence = probabilities[0][cnn_prediction].item() |
| |
| cnn_time = time.time() - start_cnn |
|
|
| |
| |
| |
|
|
| class_names = ['NORMAL', 'PNEUMONIA'] |
|
|
| response = { |
| "success": True, |
| "predictions": { |
| "svm": { |
| "class": class_names[svm_prediction], |
| "class_id": int(svm_prediction), |
| "probabilities": { |
| "NORMAL": float(svm_proba[0]), |
| "PNEUMONIA": float(svm_proba[1]) |
| }, |
| "confidence": float(svm_proba[svm_prediction]), |
| "inference_time_ms": round(svm_time * 1000, 2) |
| }, |
| "cnn": { |
| "class": class_names[cnn_prediction], |
| "class_id": int(cnn_prediction), |
| "confidence": float(cnn_confidence), |
| "probabilities": { |
| "NORMAL": float(probabilities[0][0]), |
| "PNEUMONIA": float(probabilities[0][1]) |
| }, |
| "inference_time_ms": round(cnn_time * 1000, 2) |
| } |
| }, |
| "consensus": { |
| "agreement": bool(svm_prediction == cnn_prediction), |
| "predicted_class": class_names[cnn_prediction] if svm_prediction == cnn_prediction else "DISAGREEMENT" |
| }, |
| "total_time_ms": round((svm_time + cnn_time) * 1000, 2), |
| "image_info": { |
| "filename": file.filename, |
| "original_size": image.size, |
| "format": image.format |
| } |
| } |
| return JSONResponse(content=response) |
| |
| except Exception as e: |
| print(f"Error during prediction: {str(e)}") |
| raise HTTPException( |
| status_code=500, |
| detail=f"Prediction failed: {str(e)}" |
| ) |
|
|
|
|
| if __name__ == "__main__": |
| |
| uvicorn.run( |
| "app:app", |
| host="0.0.0.0", |
| port=7860, |
| reload=True, |
| log_level="info" |
| ) |