NeuroHealth / app.py
Mikecode123's picture
Update app.py
8b65e00 verified
Raw
History Blame
3.13 kB
import torch
import torch.nn as nn
import torchvision.models as models
from fastapi import FastAPI, UploadFile, File
from PIL import Image
import io
import torchvision.transforms as transforms
from pathlib import Path
app = FastAPI(title="Alzheimer Ensemble API")
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# -----------------------------
# LABELS (YOUR REAL ONES)
# -----------------------------
LABELS = [
"Mild Demented",
"Moderate Demented",
"Non Demented",
"Very Mild Demented"
]
# -----------------------------
# IMAGE PREPROCESSING
# -----------------------------
transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5])
])
# -----------------------------
# MODEL ARCHITECTURE (4-LAYER HEAD)
# -----------------------------
def build_model():
model = models.densenet121(weights=None)
model.classifier = nn.Sequential(
nn.Linear(1024, 512),
nn.ReLU(),
nn.Dropout(0.3),
nn.Linear(512, 4)
)
return model
# -----------------------------
# LOAD MODEL
# -----------------------------
def load_model(path):
model = build_model()
state = torch.load(path, map_location=DEVICE)
# safe load for HF spaces
model.load_state_dict(state, strict=False)
model.to(DEVICE)
model.eval()
return model
# -----------------------------
# MODEL PATHS
# -----------------------------
BASE = Path("saved_models")
model_121 = load_model(BASE / "alzheimers_densenet121.pth")
model_169 = load_model(BASE / "alzheimers_densenet169.pth")
model_201 = load_model(BASE / "alzheimers_densenet201.pth")
# -----------------------------
# SINGLE MODEL PREDICTION
# -----------------------------
def predict_single(model, image_tensor):
with torch.no_grad():
out = model(image_tensor)
probs = torch.softmax(out, dim=1)
return probs[0]
# -----------------------------
# ENSEMBLE
# -----------------------------
def ensemble_predict(image_bytes):
image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
image = transform(image).unsqueeze(0).to(DEVICE)
p1 = predict_single(model_121, image)
p2 = predict_single(model_169, image)
p3 = predict_single(model_201, image)
avg = (p1 + p2 + p3) / 3
confidence, cls = torch.max(avg, dim=0)
return {
"prediction": LABELS[cls.item()],
"class_id": cls.item(),
"confidence": round(confidence.item() * 100, 2),
"probabilities": {
LABELS[i]: round(avg[i].item() * 100, 2)
for i in range(len(LABELS))
}
}
# -----------------------------
# ROUTES
# -----------------------------
@app.get("/")
def home():
return {
"message": "Alzheimer Ensemble API Running",
"classes": LABELS
}
@app.post("/predict")
async def predict(file: UploadFile = File(...)):
image_bytes = await file.read()
return ensemble_predict(image_bytes)