Mikecode123 commited on
Commit
89cfd03
·
verified ·
1 Parent(s): 4675671

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +26 -53
app.py CHANGED
@@ -4,7 +4,7 @@ import torch
4
  import torch.nn as nn
5
  import numpy as np
6
 
7
- app = FastAPI(title="NeuroHealth EEG API", version="RESET-2")
8
 
9
  DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
10
 
@@ -12,7 +12,6 @@ DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
12
  # LABELS
13
  # ======================
14
  AD_CLASSES = ["Alzheimer", "FTD", "Control"]
15
- PD_CLASSES = ["Parkinson", "Control"]
16
 
17
  # ======================
18
  # INPUT MODEL
@@ -28,35 +27,15 @@ class EEG_CNN_AD(nn.Module):
28
  super().__init__()
29
  self.conv1 = nn.Conv1d(19, 32, 7, padding=3)
30
  self.bn1 = nn.BatchNorm1d(32)
31
- self.conv2 = nn.Conv1d(32, 64, 5, padding=2)
32
- self.bn2 = nn.BatchNorm1d(64)
33
- self.conv3 = nn.Conv1d(64, 128, 3, padding=1)
34
- self.bn3 = nn.BatchNorm1d(128)
35
- self.pool = nn.AdaptiveAvgPool1d(1)
36
- self.fc = nn.Linear(128, 3)
37
 
38
- def forward(self, x):
39
- x = x.view(x.size(0), 19, 76)
40
- x = torch.relu(self.bn1(self.conv1(x)))
41
- x = torch.relu(self.bn2(self.conv2(x)))
42
- x = torch.relu(self.bn3(self.conv3(x)))
43
- x = self.pool(x).squeeze(-1)
44
- return self.fc(x)
45
-
46
- # ======================
47
- # PD MODEL - 1D CNN (Fixed to match checkpoint)
48
- # ======================
49
- class EEG_CNN_PD(nn.Module):
50
- def __init__(self):
51
- super().__init__()
52
- self.conv1 = nn.Conv1d(19, 32, 7, padding=3)
53
- self.bn1 = nn.BatchNorm1d(32)
54
  self.conv2 = nn.Conv1d(32, 64, 5, padding=2)
55
  self.bn2 = nn.BatchNorm1d(64)
 
56
  self.conv3 = nn.Conv1d(64, 128, 3, padding=1)
57
  self.bn3 = nn.BatchNorm1d(128)
 
58
  self.pool = nn.AdaptiveAvgPool1d(1)
59
- self.fc = nn.Linear(128, 2)
60
 
61
  def forward(self, x):
62
  x = x.view(x.size(0), 19, 76)
@@ -67,36 +46,32 @@ class EEG_CNN_PD(nn.Module):
67
  return self.fc(x)
68
 
69
  # ======================
70
- # LOAD MODELS
71
  # ======================
72
  AD_PATH = "AD_eeg_cnn_ad_ftd_cn.pt"
73
- PD_PATH = "PD_CNN_FINAL.pth"
74
 
75
  print("Loading AD model...")
76
  ad_model = EEG_CNN_AD().to(DEVICE)
77
- ad_model.load_state_dict(torch.load(AD_PATH, map_location=DEVICE))
78
- ad_model.eval()
79
- print("AD loaded successfully")
80
 
81
- print("Loading PD model...")
82
- pd_model = EEG_CNN_PD().to(DEVICE)
83
- pd_model.load_state_dict(torch.load(PD_PATH, map_location=DEVICE))
84
- pd_model.eval()
85
- print("PD loaded successfully")
 
86
 
87
  # ======================
88
  # PREPROCESS
89
  # ======================
90
- def preprocess(features, mode):
91
  x = torch.tensor(features, dtype=torch.float32).to(DEVICE)
92
- if mode == "ad" or mode == "pd":
93
- # Ensure shape is (batch, channels, time) -> (1, 19, 76)
94
- if len(x.shape) == 1:
95
- x = x.view(1, 19, 76)
96
- elif len(x.shape) == 2:
97
- x = x.unsqueeze(0)
98
- return x
99
- raise ValueError("Unknown mode")
100
 
101
  # ======================
102
  # PREDICT FUNCTION
@@ -106,10 +81,13 @@ def predict(model, x, classes):
106
  out = model(x)
107
  probs = torch.softmax(out, dim=1).cpu().numpy()[0]
108
  pred = int(np.argmax(probs))
 
109
  return {
110
  "prediction": classes[pred],
111
  "confidence": float(probs[pred]),
112
- "probabilities": {classes[i]: float(probs[i]) for i in range(len(classes))}
 
 
113
  }
114
 
115
  # ======================
@@ -119,16 +97,11 @@ def predict(model, x, classes):
119
  def home():
120
  return {
121
  "status": "ready",
122
- "message": "NeuroHealth EEG API is running (Fixed CNN architecture)",
123
  "year": 2026
124
  }
125
 
126
  @app.post("/predict/ad")
127
  def predict_ad(req: EEGRequest):
128
- x = preprocess(req.features, "ad")
129
- return predict(ad_model, x, AD_CLASSES)
130
-
131
- @app.post("/predict/pd")
132
- def predict_pd(req: EEGRequest):
133
- x = preprocess(req.features, "pd")
134
- return predict(pd_model, x, PD_CLASSES)
 
4
  import torch.nn as nn
5
  import numpy as np
6
 
7
+ app = FastAPI(title="NeuroHealth EEG API", version="DEMO-AD-ONLY")
8
 
9
  DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
10
 
 
12
  # LABELS
13
  # ======================
14
  AD_CLASSES = ["Alzheimer", "FTD", "Control"]
 
15
 
16
  # ======================
17
  # INPUT MODEL
 
27
  super().__init__()
28
  self.conv1 = nn.Conv1d(19, 32, 7, padding=3)
29
  self.bn1 = nn.BatchNorm1d(32)
 
 
 
 
 
 
30
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
31
  self.conv2 = nn.Conv1d(32, 64, 5, padding=2)
32
  self.bn2 = nn.BatchNorm1d(64)
33
+
34
  self.conv3 = nn.Conv1d(64, 128, 3, padding=1)
35
  self.bn3 = nn.BatchNorm1d(128)
36
+
37
  self.pool = nn.AdaptiveAvgPool1d(1)
38
+ self.fc = nn.Linear(128, 3)
39
 
40
  def forward(self, x):
41
  x = x.view(x.size(0), 19, 76)
 
46
  return self.fc(x)
47
 
48
  # ======================
49
+ # LOAD MODEL
50
  # ======================
51
  AD_PATH = "AD_eeg_cnn_ad_ftd_cn.pt"
 
52
 
53
  print("Loading AD model...")
54
  ad_model = EEG_CNN_AD().to(DEVICE)
 
 
 
55
 
56
+ ad_model.load_state_dict(
57
+ torch.load(AD_PATH, map_location=DEVICE)
58
+ )
59
+
60
+ ad_model.eval()
61
+ print("AD model loaded successfully")
62
 
63
  # ======================
64
  # PREPROCESS
65
  # ======================
66
+ def preprocess(features):
67
  x = torch.tensor(features, dtype=torch.float32).to(DEVICE)
68
+
69
+ if len(x.shape) == 1:
70
+ x = x.view(1, 19, 76)
71
+ elif len(x.shape) == 2:
72
+ x = x.unsqueeze(0)
73
+
74
+ return x
 
75
 
76
  # ======================
77
  # PREDICT FUNCTION
 
81
  out = model(x)
82
  probs = torch.softmax(out, dim=1).cpu().numpy()[0]
83
  pred = int(np.argmax(probs))
84
+
85
  return {
86
  "prediction": classes[pred],
87
  "confidence": float(probs[pred]),
88
+ "probabilities": {
89
+ classes[i]: float(probs[i]) for i in range(len(classes))
90
+ }
91
  }
92
 
93
  # ======================
 
97
  def home():
98
  return {
99
  "status": "ready",
100
+ "message": "NeuroHealth EEG API - Alzheimer Demo Mode",
101
  "year": 2026
102
  }
103
 
104
  @app.post("/predict/ad")
105
  def predict_ad(req: EEGRequest):
106
+ x = preprocess(req.features)
107
+ return predict(ad_model, x, AD_CLASSES)