Spaces:
Sleeping
Sleeping
| import torch | |
| import torch.nn as nn | |
| import torchvision.models as models | |
| from fastapi import FastAPI, UploadFile, File | |
| from fastapi.middleware.cors import CORSMiddleware | |
| from fastapi.responses import JSONResponse | |
| from PIL import Image | |
| import io | |
| import torchvision.transforms as transforms | |
| # ========================= | |
| # APP INIT | |
| # ========================= | |
| app = FastAPI(title="Alzheimer Ensemble API", version="1.0") | |
| # ========================= | |
| # CORS | |
| # ========================= | |
| app.add_middleware( | |
| CORSMiddleware, | |
| allow_origins=["*"], | |
| allow_credentials=True, | |
| allow_methods=["*"], | |
| allow_headers=["*"], | |
| ) | |
| # ========================= | |
| # DEVICE | |
| # ========================= | |
| DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| print("Using device:", DEVICE) | |
| # ========================= | |
| # CLASS LABELS (FIXED) | |
| # ========================= | |
| CLASSES = [ | |
| "Mild Demented", | |
| "Moderate Demented", | |
| "Non Demented", | |
| "Very Mild Demented" | |
| ] | |
| # ========================= | |
| # IMAGE TRANSFORM | |
| # ========================= | |
| transform = transforms.Compose([ | |
| transforms.Resize((224, 224)), | |
| transforms.ToTensor(), | |
| transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5]) | |
| ]) | |
| # ========================= | |
| # MODEL BUILDER | |
| # ========================= | |
| def build_model(version="121"): | |
| if version == "121": | |
| model = models.densenet121(weights=None) | |
| in_features = 1024 | |
| elif version == "169": | |
| model = models.densenet169(weights=None) | |
| in_features = 1664 | |
| else: | |
| model = models.densenet201(weights=None) | |
| in_features = 1920 | |
| model.classifier = nn.Sequential( | |
| nn.Dropout(0.4), | |
| nn.Linear(in_features, len(CLASSES)) | |
| ) | |
| return model | |
| # ========================= | |
| # SAFE MODEL LOADER | |
| # ========================= | |
| def load_model(path, version): | |
| model = build_model(version) | |
| try: | |
| checkpoint = torch.load(path, map_location=DEVICE) | |
| if isinstance(checkpoint, dict): | |
| if "state_dict" in checkpoint: | |
| checkpoint = checkpoint["state_dict"] | |
| elif "model_state_dict" in checkpoint: | |
| checkpoint = checkpoint["model_state_dict"] | |
| model.load_state_dict(checkpoint, strict=False) | |
| print(f"Loaded: {path}") | |
| except Exception as e: | |
| print(f"Failed loading {path}: {e}") | |
| model.to(DEVICE) | |
| model.eval() | |
| return model | |
| # ========================= | |
| # LOAD MODELS | |
| # ========================= | |
| model_121 = load_model("alzheimers_densenet121.pth", "121") | |
| model_169 = load_model("alzheimers_densenet169.pth", "169") | |
| model_201 = load_model("alzheimers_densenet201.pth", "201") | |
| # ========================= | |
| # IMAGE PROCESSING | |
| # ========================= | |
| def process_image(image_bytes): | |
| img = Image.open(io.BytesIO(image_bytes)).convert("RGB") | |
| img = transform(img).unsqueeze(0).to(DEVICE) | |
| return img | |
| # ========================= | |
| # PREDICTION FUNCTION | |
| # All confidence values are returned as floats in range [0.0, 1.0] | |
| # ========================= | |
| def predict(model, x): | |
| with torch.no_grad(): | |
| output = model(x) | |
| probs = torch.softmax(output, dim=1)[0] | |
| conf, cls = torch.max(probs, 0) | |
| cls = int(cls.item()) | |
| return { | |
| "prediction": CLASSES[cls], | |
| "class_id": cls, | |
| # Confidence as 0.0–1.0 decimal | |
| "confidence": float(conf.item()), | |
| "probabilities": { | |
| CLASSES[i]: float(probs[i].item()) | |
| for i in range(len(CLASSES)) | |
| } | |
| } | |
| # ========================= | |
| # HEALTH ENDPOINT | |
| # ========================= | |
| def health(): | |
| return { | |
| "status": "running", | |
| "service": "Alzheimer MRI Ensemble API", | |
| "models_loaded": { | |
| "densenet121": True, | |
| "densenet169": True, | |
| "densenet201": True, | |
| }, | |
| "device": str(DEVICE), | |
| "classes": CLASSES, | |
| } | |
| # ========================= | |
| # ROOT ENDPOINT | |
| # ========================= | |
| def home(): | |
| return { | |
| "status": "running", | |
| "models": ["121", "169", "201"], | |
| "classes": CLASSES, | |
| "endpoints": [ | |
| "/health", | |
| "/predict/121", | |
| "/predict/169", | |
| "/predict/201", | |
| "/predict/ensemble" | |
| ] | |
| } | |
| # ========================= | |
| # SINGLE MODEL PREDICTIONS | |
| # ========================= | |
| async def predict_121(file: UploadFile = File(...)): | |
| img = process_image(await file.read()) | |
| result = predict(model_121, img) | |
| result["model"] = "densenet121" | |
| return JSONResponse(result) | |
| async def predict_169(file: UploadFile = File(...)): | |
| img = process_image(await file.read()) | |
| result = predict(model_169, img) | |
| result["model"] = "densenet169" | |
| return JSONResponse(result) | |
| async def predict_201(file: UploadFile = File(...)): | |
| img = process_image(await file.read()) | |
| result = predict(model_201, img) | |
| result["model"] = "densenet201" | |
| return JSONResponse(result) | |
| # ========================= | |
| # ENSEMBLE PREDICTION | |
| # Returns ensemble result with individual model results | |
| # All confidence values are 0.0-1.0 | |
| # ========================= | |
| async def ensemble(file: UploadFile = File(...)): | |
| img = process_image(await file.read()) | |
| r1 = predict(model_121, img) | |
| r2 = predict(model_169, img) | |
| r3 = predict(model_201, img) | |
| # Average probabilities (all already 0.0-1.0) | |
| avg_probs = {} | |
| for c in CLASSES: | |
| avg_probs[c] = ( | |
| r1["probabilities"][c] + | |
| r2["probabilities"][c] + | |
| r3["probabilities"][c] | |
| ) / 3 | |
| final_class = max(avg_probs, key=avg_probs.get) | |
| # Ensemble confidence = max averaged probability (0.0-1.0) | |
| ensemble_confidence = avg_probs[final_class] | |
| return JSONResponse({ | |
| "prediction": final_class, | |
| # confidence as 0.0-1.0 | |
| "confidence": ensemble_confidence, | |
| "probabilities": avg_probs, | |
| "individual_models": { | |
| "densenet121": r1, | |
| "densenet169": r2, | |
| "densenet201": r3, | |
| } | |
| }) |