ThermoDense / server.py
sunnykiyo's picture
Fix FastAPI TemplateResponse syntax for Starlette compatibility
20f9f78
Raw
History Blame Contribute Delete
9.81 kB
# =============================================================================
# server.py β€” FastAPI Backend
# =============================================================================
# This connects the beautiful HTML frontend to your trained ThermoCAN model.
#
# HOW TO RUN:
# pip install -r requirements.txt
# python server.py
#
# Then open browser: http://localhost:8000
# =============================================================================
import os
import sys
import io
import base64
import yaml
import torch
import numpy as np
import pandas as pd
from pathlib import Path
from PIL import Image
from fastapi import FastAPI, File, UploadFile, Form, HTTPException
from fastapi.responses import HTMLResponse, JSONResponse
from fastapi.staticfiles import StaticFiles
from fastapi.templating import Jinja2Templates
from fastapi import Request
from contextlib import asynccontextmanager
import uvicorn
sys.path.insert(0, str(Path(__file__).parent))
from model_loader import load_model
from predictor import predict, generate_gradcam, tiff_float_to_uint8
# ── Configuration ─────────────────────────────────────────────────────────────
BASE_DIR = Path(__file__).parent.absolute()
CONFIG_PATH = BASE_DIR / "configs" / "config.yaml"
CHECKPOINT_PATH = BASE_DIR / "outputs" / "checkpoints" / "best.pth"
# ── Load config ───────────────────────────────────────────────────────────────
with open(CONFIG_PATH) as f:
CONFIG = yaml.safe_load(f)
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# Global state
MODEL = None
ENCODER = None
NUM_VIEWS = 3
MODEL_LOADED = False
@asynccontextmanager
async def lifespan(app: FastAPI):
global MODEL, ENCODER, NUM_VIEWS, MODEL_LOADED
print("Device: " + str(DEVICE))
print("Loading ThermoCAN model...")
try:
# Use str(CHECKPOINT_PATH) for os.path operations inside load_model
MODEL, ENCODER, NUM_VIEWS = load_model(str(CHECKPOINT_PATH), CONFIG, DEVICE)
MODEL_LOADED = True
print("Model loaded and ready!")
except Exception as e:
print(f"Model load error: {e}")
MODEL_LOADED = False
yield
# Cleanup if needed
# Which views to use
d = CONFIG["dataset"]
use_all = d.get("use_all_five_views", False)
primary = d.get("primary_views", [0, 1, 3])
WANTED_VIEWS = [0, 1, 2, 3, 4] if use_all else primary
IMAGE_SIZE = d.get("image_size", 224)
# ── FastAPI app ───────────────────────────────────────────────────────────────
app = FastAPI(title="ThermoCAN", version="1.0", lifespan=lifespan)
# Serve HTML templates
templates = Jinja2Templates(directory="templates")
@app.get("/", response_class=HTMLResponse)
async def home(request: Request):
"""Serve the main frontend page."""
return templates.TemplateResponse(request=request, name="index.html")
@app.get("/favicon.ico")
async def favicon():
"""Handle favicon requests."""
return JSONResponse(status_code=204, content={})
@app.get("/health")
async def health():
"""Health check β€” tells you if model is loaded."""
return {
"status": "ok" if MODEL_LOADED else "model_not_loaded",
"device": str(DEVICE),
"num_views": NUM_VIEWS,
"wanted_views": WANTED_VIEWS,
"clinical_dim": ENCODER.feature_dim if ENCODER else 0,
}
@app.post("/predict")
async def predict_endpoint(
# Image uploads β€” view_0 through view_4
view_0: UploadFile = File(None),
view_1: UploadFile = File(None),
view_2: UploadFile = File(None),
view_3: UploadFile = File(None),
view_4: UploadFile = File(None),
# Clinical fields
age: str = Form("45"),
body_temp: str = Form("36.6"),
biopsy: str = Form("No"),
mammography: str = Form("No"),
radiotherapy: str = Form("No"),
mastectomy: str = Form("No"),
eating_habits: str = Form("Low fat"),
protocol_smoked: str = Form("No"),
marital_status: str = Form("Married"),
use_tta: str = Form("true"),
):
"""
Main prediction endpoint.
Receives images + clinical data, returns prediction + Grad-CAM.
"""
if not MODEL_LOADED:
raise HTTPException(
status_code=503,
detail="Model not loaded. Check CHECKPOINT_PATH in server.py."
)
# ── Collect uploaded images ────────────────────────────────────────────
uploads = {0: view_0, 1: view_1, 2: view_2, 3: view_3, 4: view_4}
images_dict = {}
for view_id, upload in uploads.items():
if upload is not None and upload.filename:
try:
raw = await upload.read()
img = Image.open(io.BytesIO(raw))
# Convert mode F (TIFF float) or any mode to RGB PIL image
if img.mode == "F":
arr = tiff_float_to_uint8(img)
img = Image.fromarray(arr)
elif img.mode != "RGB":
img = img.convert("RGB")
images_dict[view_id] = img
except Exception as e:
print(f"Failed to read view_{view_id}: {e}")
if not images_dict:
raise HTTPException(
status_code=400,
detail="No valid images uploaded. Please upload at least the Frontal view."
)
if 0 not in images_dict:
raise HTTPException(
status_code=400,
detail="Frontal view (view_0) is required but was not uploaded."
)
# ── Encode clinical features ───────────────────────────────────────────
clinical_tensor = None
if ENCODER is not None and ENCODER.feature_dim > 0:
try:
row = {
"age_current": float(age),
"age_at_visit": float(age),
"body_temperature": float(body_temp),
"menarche": 12.0,
"biopsy": biopsy,
"mammography": mammography,
"radiotherapy": radiotherapy,
"use_of_hormone_replacement": "No",
"plastic_surgery": "No",
"prosthesis": "No",
"marital_status": marital_status,
"eating_habits": eating_habits,
"protocol_smoked": protocol_smoked,
"protocol_drank_coffee": "No",
"mastectomy": mastectomy,
}
df = pd.DataFrame([row])
arr = ENCODER.transform(df)
clinical_tensor = torch.tensor(arr[0], dtype=torch.float32)
except Exception as e:
print(f"Clinical encoding failed: {e} β€” using image only")
clinical_tensor = None
# ── Run prediction ─────────────────────────────────────────────────────
use_tta_bool = use_tta.lower() in ("true", "1", "yes")
try:
probability, label_str, confidence = predict(
MODEL, images_dict, clinical_tensor, DEVICE,
NUM_VIEWS, WANTED_VIEWS,
use_tta=use_tta_bool,
image_size=IMAGE_SIZE,
)
except Exception as e:
raise HTTPException(status_code=500, detail=f"Prediction error: {str(e)}")
# ── Generate Grad-CAM ──────────────────────────────────────────────────
gradcam_b64 = None
try:
gradcam_pil = generate_gradcam(
MODEL, images_dict, WANTED_VIEWS, DEVICE, IMAGE_SIZE
)
buf = io.BytesIO()
gradcam_pil.save(buf, format="PNG")
gradcam_b64 = base64.b64encode(buf.getvalue()).decode("utf-8")
except Exception as e:
print(f" ⚠️ Grad-CAM failed: {e}")
# ── Clean label for frontend ───────────────────────────────────────────
# Remove emoji from label string for JSON
clean_label = "MALIGNANT" if probability >= 0.5 else "BENIGN"
return JSONResponse({
"probability": round(float(probability), 4),
"label": clean_label,
"confidence": confidence,
"gradcam_base64": gradcam_b64,
"views_used": len(images_dict),
"tta_used": use_tta_bool,
"clinical_used": clinical_tensor is not None,
"device": str(DEVICE),
})
# ── Run server ────────────────────────────────────────────────────────────────
if __name__ == "__main__":
print("=" * 55)
print("ThermoCAN -- Starting Web Server")
print("=" * 55)
print(f" Open browser: http://localhost:8000")
print(f" Press Ctrl+C to stop")
print("=" * 55)
uvicorn.run(
app,
host="0.0.0.0",
port=7860,
reload=False,
log_level="info",
)