Mikecode123 commited on
Commit
720d7d5
·
verified ·
1 Parent(s): c8092e3

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +328 -115
app.py CHANGED
@@ -1,153 +1,366 @@
 
 
 
 
 
 
 
1
  import torch
2
  import torch.nn as nn
3
  import torchvision.models as models
 
4
  from fastapi import FastAPI, UploadFile, File
5
- from PIL import Image
6
- import io
7
- import torchvision.transforms as transforms
8
- import numpy as np
9
- import tensorflow as tf
10
 
11
- app = FastAPI(title="Alzheimer Ensemble API")
 
 
 
12
 
13
- DEVICE = torch.device("cpu")
 
 
14
 
15
- # =========================
16
- # LABELS
17
- # =========================
18
  LABELS = [
19
- "Mild Demented",
20
- "Moderate Demented",
21
- "Non Demented",
22
- "Very Mild Demented"
23
  ]
24
 
25
- # =========================
26
- # TRANSFORM
27
- # =========================
28
- transform = transforms.Compose([
29
- transforms.Resize((224, 224)),
30
- transforms.ToTensor(),
31
- transforms.Normalize([0.5]*3, [0.5]*3)
32
- ])
33
-
34
- # =========================
35
- # SAFE KERAS LOADER
36
- # =========================
37
- def load_keras_model(path):
38
- try:
39
- return tf.keras.models.load_model(path, compile=False)
40
- except Exception as e:
41
- print("Keras load failed:", path, e)
42
- return None
43
-
44
- # =========================
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
45
  # LOAD MODELS
46
- # =========================
47
- model_121 = load_keras_model("densenet121_parkinsonsDATSCAN.keras")
48
- model_169 = load_keras_model("parkinsons_densenet169DATSCAN.keras")
49
- model_201 = load_keras_model("parkinsons_densenet201DATSCAN.keras")
 
50
 
51
- models_list = [model_121, model_169, model_201]
 
52
 
53
- # =========================
54
- # SINGLE PREDICTION (KERAS SAFE)
55
- # =========================
56
- def predict_single(model, image_tensor):
57
- if model is None:
58
- return np.zeros(len(LABELS))
59
 
60
- img = image_tensor.permute(0, 2, 3, 1).numpy()
61
 
62
- preds = model.predict(img, verbose=0)[0]
63
- return preds
64
 
65
- # =========================
66
- # ENSEMBLE PREDICTION
67
- # =========================
68
- def ensemble_predict(image_bytes):
69
- image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
70
- image = transform(image).unsqueeze(0)
71
 
72
- preds_all = []
73
- model_confidence_report = []
74
 
75
- for i, m in enumerate(models_list):
76
- preds = predict_single(m, image)
77
- preds_all.append(preds)
 
 
78
 
79
- model_confidence_report.append({
80
- "model": f"model_{i+1}",
81
- "confidence": float(np.max(preds)),
82
- "prediction": LABELS[int(np.argmax(preds))]
83
- })
84
 
85
- avg = np.mean(preds_all, axis=0)
86
 
87
- cls = int(np.argmax(avg))
88
- conf = float(np.max(avg))
89
 
90
  return {
91
- "prediction": LABELS[cls],
92
- "class_id": cls,
93
- "confidence": round(conf * 100, 2),
94
  "probabilities": {
95
- LABELS[i]: round(float(avg[i]) * 100, 2)
96
- for i in range(len(LABELS))
97
- },
98
- "model_breakdown": model_confidence_report
 
 
 
 
 
 
 
 
 
 
99
  }
100
 
101
- # =========================
 
102
  # INDIVIDUAL ENDPOINTS
103
- # =========================
104
- @app.post("/predict/121")
105
  async def predict_121(file: UploadFile = File(...)):
106
- img = await file.read()
107
- return ensemble_predict_single(img, model_121)
108
 
109
- @app.post("/predict/169")
 
 
 
 
 
 
 
 
110
  async def predict_169(file: UploadFile = File(...)):
111
- img = await file.read()
112
- return ensemble_predict_single(img, model_169)
113
 
114
- @app.post("/predict/201")
 
 
 
 
 
 
 
 
115
  async def predict_201(file: UploadFile = File(...)):
116
- img = await file.read()
117
- return ensemble_predict_single(img, model_201)
118
 
119
- # helper
120
- def ensemble_predict_single(image_bytes, model):
121
- image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
122
- image = transform(image).unsqueeze(0)
123
 
124
- preds = predict_single(model, image)
125
- cls = int(np.argmax(preds))
126
 
127
- return {
128
- "prediction": LABELS[cls],
129
- "confidence": float(np.max(preds)),
130
- "probabilities": {
131
- LABELS[i]: float(preds[i])
132
- for i in range(len(LABELS))
133
- }
134
- }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
135
 
136
- # =========================
137
- # ENSEMBLE ENDPOINT
138
- # =========================
139
- @app.post("/predict")
140
- async def predict(file: UploadFile = File(...)):
141
- image_bytes = await file.read()
142
- return ensemble_predict(image_bytes)
143
-
144
- # =========================
145
- # HEALTH CHECK
146
- # =========================
147
- @app.get("/")
148
- def home():
149
  return {
150
- "status": "running",
151
- "models": ["121", "169", "201"],
152
- "ensemble": True
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
153
  }
 
1
+ # =========================================
2
+ # IMPORTS
3
+ # =========================================
4
+ import io
5
+ import cv2
6
+ import nibabel as nib
7
+ import numpy as np
8
  import torch
9
  import torch.nn as nn
10
  import torchvision.models as models
11
+
12
  from fastapi import FastAPI, UploadFile, File
13
+ from fastapi.responses import JSONResponse
 
 
 
 
14
 
15
+ # =========================================
16
+ # CONFIG
17
+ # =========================================
18
+ DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
19
 
20
+ IMG_SIZE = 128
21
+ NUM_CLASSES = 3
22
+ NUM_SLICES = 32
23
 
 
 
 
24
  LABELS = [
25
+ "Control",
26
+ "Prodromal",
27
+ "Parkinsons"
 
28
  ]
29
 
30
+ print("Using device:", DEVICE)
31
+
32
+ # =========================================
33
+ # FASTAPI
34
+ # =========================================
35
+ app = FastAPI(
36
+ title="Parkinsons DATSCAN Ensemble API",
37
+ version="1.0"
38
+ )
39
+
40
+ # =========================================
41
+ # PREPROCESSING
42
+ # =========================================
43
+ def load_nifti_from_bytes(file_bytes):
44
+ temp_path = "temp.nii"
45
+
46
+ with open(temp_path, "wb") as f:
47
+ f.write(file_bytes)
48
+
49
+ volume = nib.load(temp_path).get_fdata()
50
+ volume = np.squeeze(volume)
51
+
52
+ return volume
53
+
54
+
55
+ def preprocess_2d(volume):
56
+ depth = volume.shape[2]
57
+
58
+ idx1 = np.linspace(0, depth//3 - 1, 10).astype(int)
59
+ idx2 = np.linspace(depth//3, 2*depth//3 - 1, 10).astype(int)
60
+ idx3 = np.linspace(2*depth//3, depth - 1, 12).astype(int)
61
+
62
+ def make_channel(indices):
63
+ slices = []
64
+
65
+ for i in indices:
66
+ sl = volume[:, :, i]
67
+
68
+ sl = sl - sl.min()
69
+ sl = sl / (sl.max() + 1e-6)
70
+
71
+ sl = cv2.resize(sl, (IMG_SIZE, IMG_SIZE))
72
+ slices.append(sl)
73
+
74
+ return np.mean(slices, axis=0)
75
+
76
+ r = make_channel(idx1)
77
+ g = make_channel(idx2)
78
+ b = make_channel(idx3)
79
+
80
+ img = np.stack([r, g, b], axis=0)
81
+
82
+ return torch.tensor(img, dtype=torch.float32).unsqueeze(0)
83
+
84
+
85
+ def preprocess_3d(volume):
86
+ depth = volume.shape[2]
87
+
88
+ indices = np.linspace(0, depth - 1, NUM_SLICES).astype(int)
89
+
90
+ slices = []
91
+
92
+ for i in indices:
93
+ sl = volume[:, :, i]
94
+
95
+ sl = sl - sl.min()
96
+ sl = sl / (sl.max() + 1e-6)
97
+
98
+ sl = cv2.resize(sl, (IMG_SIZE, IMG_SIZE))
99
+ slices.append(sl)
100
+
101
+ vol = np.stack(slices)
102
+
103
+ vol = torch.tensor(vol, dtype=torch.float32)
104
+
105
+ vol = vol.unsqueeze(0).unsqueeze(0)
106
+
107
+ return vol
108
+
109
+
110
+ # =========================================
111
+ # DENSENET MODELS
112
+ # =========================================
113
+ class DenseNet121Model(nn.Module):
114
+ def __init__(self):
115
+ super().__init__()
116
+
117
+ self.base = models.densenet121(weights=None)
118
+
119
+ self.base.features.conv0 = nn.Conv2d(
120
+ 3,
121
+ 64,
122
+ kernel_size=7,
123
+ stride=2,
124
+ padding=3,
125
+ bias=False
126
+ )
127
+
128
+ self.base.classifier = nn.Linear(
129
+ self.base.classifier.in_features,
130
+ NUM_CLASSES
131
+ )
132
+
133
+ def forward(self, x):
134
+ return self.base(x)
135
+
136
+
137
+ class DenseNet169Model(nn.Module):
138
+ def __init__(self):
139
+ super().__init__()
140
+
141
+ self.base = models.densenet169(weights=None)
142
+
143
+ self.base.features.conv0 = nn.Conv2d(
144
+ 3,
145
+ 64,
146
+ kernel_size=7,
147
+ stride=2,
148
+ padding=3,
149
+ bias=False
150
+ )
151
+
152
+ self.base.classifier = nn.Linear(
153
+ self.base.classifier.in_features,
154
+ NUM_CLASSES
155
+ )
156
+
157
+ def forward(self, x):
158
+ return self.base(x)
159
+
160
+
161
+ class DenseNet201Model(nn.Module):
162
+ def __init__(self):
163
+ super().__init__()
164
+
165
+ self.base = models.densenet201(weights=None)
166
+
167
+ self.base.features.conv0 = nn.Conv2d(
168
+ 3,
169
+ 64,
170
+ kernel_size=7,
171
+ stride=2,
172
+ padding=3,
173
+ bias=False
174
+ )
175
+
176
+ self.base.classifier = nn.Linear(
177
+ self.base.classifier.in_features,
178
+ NUM_CLASSES
179
+ )
180
+
181
+ def forward(self, x):
182
+ return self.base(x)
183
+
184
+
185
+ # =========================================
186
+ # 3D CNN
187
+ # =========================================
188
+ class CNN3D(nn.Module):
189
+ def __init__(self):
190
+ super().__init__()
191
+
192
+ self.net = nn.Sequential(
193
+
194
+ nn.Conv3d(1, 16, 3, padding=1),
195
+ nn.ReLU(),
196
+ nn.MaxPool3d(2),
197
+
198
+ nn.Conv3d(16, 32, 3, padding=1),
199
+ nn.ReLU(),
200
+ nn.MaxPool3d(2),
201
+
202
+ nn.Conv3d(32, 64, 3, padding=1),
203
+ nn.ReLU(),
204
+ nn.MaxPool3d(2)
205
+
206
+ )
207
+
208
+ self.fc = nn.Sequential(
209
+ nn.Linear(64 * 4 * 16 * 16, 256),
210
+ nn.ReLU(),
211
+ nn.Dropout(0.3),
212
+ nn.Linear(256, NUM_CLASSES)
213
+ )
214
+
215
+ def forward(self, x):
216
+ x = self.net(x)
217
+ x = x.view(x.size(0), -1)
218
+ return self.fc(x)
219
+
220
+
221
+ # =========================================
222
  # LOAD MODELS
223
+ # =========================================
224
+ def load_model(model, path):
225
+ model.load_state_dict(
226
+ torch.load(path, map_location=DEVICE)
227
+ )
228
 
229
+ model.to(DEVICE)
230
+ model.eval()
231
 
232
+ print(f"Loaded: {path}")
 
 
 
 
 
233
 
234
+ return model
235
 
 
 
236
 
237
+ model121 = load_model(DenseNet121Model(), "densenet121.pth")
238
+ model169 = load_model(DenseNet169Model(), "densenet169.pth")
239
+ model201 = load_model(DenseNet201Model(), "densenet201.pth")
240
+ model3d = load_model(CNN3D(), "cnn3d.pth")
 
 
241
 
 
 
242
 
243
+ # =========================================
244
+ # PREDICTION HELPERS
245
+ # =========================================
246
+ def predict_model(model, tensor):
247
+ tensor = tensor.to(DEVICE)
248
 
249
+ with torch.no_grad():
250
+ out = model(tensor)
 
 
 
251
 
252
+ probs = torch.softmax(out, dim=1)[0]
253
 
254
+ pred_idx = torch.argmax(probs).item()
 
255
 
256
  return {
257
+ "prediction": LABELS[pred_idx],
258
+ "class_id": pred_idx,
259
+ "confidence": round(float(probs[pred_idx]) * 100, 2),
260
  "probabilities": {
261
+ LABELS[i]: round(float(probs[i]) * 100, 2)
262
+ for i in range(NUM_CLASSES)
263
+ }
264
+ }, probs
265
+
266
+
267
+ # =========================================
268
+ # ROOT
269
+ # =========================================
270
+ @app.get("/")
271
+ def home():
272
+ return {
273
+ "message": "Parkinson DATSCAN Ensemble API Running",
274
+ "classes": LABELS
275
  }
276
 
277
+
278
+ # =========================================
279
  # INDIVIDUAL ENDPOINTS
280
+ # =========================================
281
+ @app.post("/predict/densenet121")
282
  async def predict_121(file: UploadFile = File(...)):
 
 
283
 
284
+ volume = load_nifti_from_bytes(await file.read())
285
+ tensor = preprocess_2d(volume)
286
+
287
+ result, _ = predict_model(model121, tensor)
288
+
289
+ return JSONResponse(result)
290
+
291
+
292
+ @app.post("/predict/densenet169")
293
  async def predict_169(file: UploadFile = File(...)):
 
 
294
 
295
+ volume = load_nifti_from_bytes(await file.read())
296
+ tensor = preprocess_2d(volume)
297
+
298
+ result, _ = predict_model(model169, tensor)
299
+
300
+ return JSONResponse(result)
301
+
302
+
303
+ @app.post("/predict/densenet201")
304
  async def predict_201(file: UploadFile = File(...)):
 
 
305
 
306
+ volume = load_nifti_from_bytes(await file.read())
307
+ tensor = preprocess_2d(volume)
 
 
308
 
309
+ result, _ = predict_model(model201, tensor)
 
310
 
311
+ return JSONResponse(result)
312
+
313
+
314
+ @app.post("/predict/cnn3d")
315
+ async def predict_3d(file: UploadFile = File(...)):
316
+
317
+ volume = load_nifti_from_bytes(await file.read())
318
+ tensor = preprocess_3d(volume)
319
+
320
+ result, _ = predict_model(model3d, tensor)
321
+
322
+ return JSONResponse(result)
323
+
324
+
325
+ # =========================================
326
+ # ENSEMBLE
327
+ # =========================================
328
+ @app.post("/predict/ensemble")
329
+ async def predict_ensemble(file: UploadFile = File(...)):
330
+
331
+ volume = load_nifti_from_bytes(await file.read())
332
+
333
+ tensor2d = preprocess_2d(volume)
334
+ tensor3d = preprocess_3d(volume)
335
+
336
+ r121, p121 = predict_model(model121, tensor2d)
337
+ r169, p169 = predict_model(model169, tensor2d)
338
+ r201, p201 = predict_model(model201, tensor2d)
339
+ r3d, p3d = predict_model(model3d, tensor3d)
340
+
341
+ avg_probs = (p121 + p169 + p201 + p3d) / 4
342
+
343
+ pred_idx = torch.argmax(avg_probs).item()
344
 
 
 
 
 
 
 
 
 
 
 
 
 
 
345
  return {
346
+
347
+ "ensemble_prediction": LABELS[pred_idx],
348
+
349
+ "ensemble_confidence": round(
350
+ float(avg_probs[pred_idx]) * 100,
351
+ 2
352
+ ),
353
+
354
+ "ensemble_probabilities": {
355
+ LABELS[i]: round(float(avg_probs[i]) * 100, 2)
356
+ for i in range(NUM_CLASSES)
357
+ },
358
+
359
+ "individual_models": {
360
+
361
+ "DenseNet121": r121,
362
+ "DenseNet169": r169,
363
+ "DenseNet201": r201,
364
+ "CNN3D": r3d
365
+ }
366
  }