Spaces:
Sleeping
Sleeping
File size: 6,476 Bytes
eb052a3 8b65e00 559e735 eb052a3 95fecec da5cfde 95fecec 559e735 b18a7cf b5ad082 b18a7cf b5ad082 eb052a3 b18a7cf eb052a3 b5ad082 da5cfde b5ad082 95fecec eb052a3 b5ad082 8b65e00 95fecec 8b65e00 b5ad082 da5cfde b5ad082 95fecec b18a7cf 95fecec b18a7cf 95fecec b18a7cf 95fecec da5cfde 95fecec 8b65e00 b5ad082 da5cfde b5ad082 95fecec b18a7cf da5cfde b5ad082 da5cfde 8b65e00 da5cfde b18a7cf b5ad082 eb052a3 b5ad082 b18a7cf b5ad082 95fecec b5ad082 95fecec eb052a3 b5ad082 559e735 b5ad082 95fecec 8b65e00 da5cfde 95fecec da5cfde 83238b7 b18a7cf da5cfde 559e735 b18a7cf 95fecec b18a7cf 83238b7 b18a7cf 559e735 b18a7cf da5cfde b18a7cf da5cfde 95fecec da5cfde 559e735 da5cfde 83238b7 eb052a3 b5ad082 da5cfde b5ad082 83238b7 da5cfde 559e735 da5cfde eb052a3 83238b7 da5cfde 559e735 da5cfde eb052a3 83238b7 da5cfde 559e735 eb052a3 b5ad082 559e735 b5ad082 83238b7 559e735 da5cfde 95fecec da5cfde 83238b7 da5cfde 559e735 8b65e00 559e735 da5cfde 559e735 da5cfde 559e735 da5cfde 559e735 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 | 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
# =========================
@app.get("/health")
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
# =========================
@app.get("/")
def home():
return {
"status": "running",
"models": ["121", "169", "201"],
"classes": CLASSES,
"endpoints": [
"/health",
"/predict/121",
"/predict/169",
"/predict/201",
"/predict/ensemble"
]
}
# =========================
# SINGLE MODEL PREDICTIONS
# =========================
@app.post("/predict/121")
async def predict_121(file: UploadFile = File(...)):
img = process_image(await file.read())
result = predict(model_121, img)
result["model"] = "densenet121"
return JSONResponse(result)
@app.post("/predict/169")
async def predict_169(file: UploadFile = File(...)):
img = process_image(await file.read())
result = predict(model_169, img)
result["model"] = "densenet169"
return JSONResponse(result)
@app.post("/predict/201")
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
# =========================
@app.post("/predict/ensemble")
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,
}
}) |