Mikecode123 commited on
Commit
174705b
·
verified ·
1 Parent(s): 6fa7603

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +28 -41
app.py CHANGED
@@ -13,59 +13,45 @@ app = FastAPI(
13
  DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
14
 
15
  # =========================================================
16
- # CLASS NAMES
17
  # =========================================================
18
 
19
- AD_CLASSES = [
20
- "Alzheimer",
21
- "FTD",
22
- "Control"
23
- ]
24
-
25
- PD_CLASSES = [
26
- "Parkinson",
27
- "Control"
28
- ]
29
 
30
  # =========================================================
31
- # CONFIG
32
  # =========================================================
33
 
34
- INPUT_DIM = 95
35
 
36
  # =========================================================
37
- # CNN MODEL (MATCHES SAVED WEIGHTS)
38
  # =========================================================
39
 
40
  class EEG_CNN(nn.Module):
41
  def __init__(self, output_dim):
42
  super().__init__()
43
 
44
- self.conv1 = nn.Conv1d(1, 32, kernel_size=3, padding=1)
45
  self.bn1 = nn.BatchNorm1d(32)
46
 
47
- self.conv2 = nn.Conv1d(32, 64, kernel_size=3, padding=1)
48
  self.bn2 = nn.BatchNorm1d(64)
49
 
50
- self.conv3 = nn.Conv1d(64, 128, kernel_size=3, padding=1)
51
  self.bn3 = nn.BatchNorm1d(128)
52
 
53
- # Global pooling removes dependence on sequence length
54
  self.pool = nn.AdaptiveAvgPool1d(1)
55
-
56
  self.fc = nn.Linear(128, output_dim)
57
 
58
  def forward(self, x):
59
- # x shape: (batch, features)
60
- x = x.unsqueeze(1) # (batch, 1, 95)
61
-
62
  x = torch.relu(self.bn1(self.conv1(x)))
63
  x = torch.relu(self.bn2(self.conv2(x)))
64
  x = torch.relu(self.bn3(self.conv3(x)))
65
 
66
- x = self.pool(x) # (batch, 128, 1)
67
- x = x.squeeze(-1) # (batch, 128)
68
-
69
  return self.fc(x)
70
 
71
  # =========================================================
@@ -75,13 +61,8 @@ class EEG_CNN(nn.Module):
75
  ad_model = EEG_CNN(3).to(DEVICE)
76
  pd_model = EEG_CNN(2).to(DEVICE)
77
 
78
- ad_model.load_state_dict(
79
- torch.load("AD_MLP.pt", map_location=DEVICE)
80
- )
81
-
82
- pd_model.load_state_dict(
83
- torch.load("PD_MLP.pt", map_location=DEVICE)
84
- )
85
 
86
  ad_model.eval()
87
  pd_model.eval()
@@ -89,18 +70,26 @@ pd_model.eval()
89
  print("Models Loaded Successfully")
90
 
91
  # =========================================================
92
- # REQUEST MODEL
93
  # =========================================================
94
 
95
  class EEGRequest(BaseModel):
96
  features: list
97
 
98
  # =========================================================
99
- # UTILITY
100
  # =========================================================
101
 
102
  def predict_model(model, features):
103
- x = torch.tensor(features, dtype=torch.float32).unsqueeze(0).to(DEVICE)
 
 
 
 
 
 
 
 
104
 
105
  with torch.no_grad():
106
  outputs = model(x)
@@ -112,17 +101,15 @@ def predict_model(model, features):
112
  return pred, confidence, probs.tolist()
113
 
114
  # =========================================================
115
- # ROOT
116
  # =========================================================
117
 
118
  @app.get("/")
119
  def home():
120
- return {
121
- "message": "NeuroHealth EEG API Running"
122
- }
123
 
124
  # =========================================================
125
- # ALZHEIMER PREDICTION
126
  # =========================================================
127
 
128
  @app.post("/predict/alzheimer")
@@ -140,7 +127,7 @@ def predict_alzheimer(request: EEGRequest):
140
  }
141
 
142
  # =========================================================
143
- # PARKINSON PREDICTION
144
  # =========================================================
145
 
146
  @app.post("/predict/parkinson")
 
13
  DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
14
 
15
  # =========================================================
16
+ # CLASS LABELS
17
  # =========================================================
18
 
19
+ AD_CLASSES = ["Alzheimer", "FTD", "Control"]
20
+ PD_CLASSES = ["Parkinson", "Control"]
 
 
 
 
 
 
 
 
21
 
22
  # =========================================================
23
+ # EXPECTED INPUT
24
  # =========================================================
25
 
26
+ INPUT_FEATURES = 95 # must reshape to (19, 5)
27
 
28
  # =========================================================
29
+ # CNN MODEL (MATCHES YOUR CHECKPOINT EXACTLY)
30
  # =========================================================
31
 
32
  class EEG_CNN(nn.Module):
33
  def __init__(self, output_dim):
34
  super().__init__()
35
 
36
+ self.conv1 = nn.Conv1d(in_channels=19, out_channels=32, kernel_size=7)
37
  self.bn1 = nn.BatchNorm1d(32)
38
 
39
+ self.conv2 = nn.Conv1d(32, 64, kernel_size=5)
40
  self.bn2 = nn.BatchNorm1d(64)
41
 
42
+ self.conv3 = nn.Conv1d(64, 128, kernel_size=3)
43
  self.bn3 = nn.BatchNorm1d(128)
44
 
 
45
  self.pool = nn.AdaptiveAvgPool1d(1)
 
46
  self.fc = nn.Linear(128, output_dim)
47
 
48
  def forward(self, x):
 
 
 
49
  x = torch.relu(self.bn1(self.conv1(x)))
50
  x = torch.relu(self.bn2(self.conv2(x)))
51
  x = torch.relu(self.bn3(self.conv3(x)))
52
 
53
+ x = self.pool(x)
54
+ x = x.squeeze(-1)
 
55
  return self.fc(x)
56
 
57
  # =========================================================
 
61
  ad_model = EEG_CNN(3).to(DEVICE)
62
  pd_model = EEG_CNN(2).to(DEVICE)
63
 
64
+ ad_model.load_state_dict(torch.load("AD_MLP.pt", map_location=DEVICE))
65
+ pd_model.load_state_dict(torch.load("PD_MLP.pt", map_location=DEVICE))
 
 
 
 
 
66
 
67
  ad_model.eval()
68
  pd_model.eval()
 
70
  print("Models Loaded Successfully")
71
 
72
  # =========================================================
73
+ # REQUEST SCHEMA
74
  # =========================================================
75
 
76
  class EEGRequest(BaseModel):
77
  features: list
78
 
79
  # =========================================================
80
+ # CORE PREDICTION FUNCTION
81
  # =========================================================
82
 
83
  def predict_model(model, features):
84
+ x = torch.tensor(features, dtype=torch.float32)
85
+
86
+ # reshape 95 -> (19, 5)
87
+ if x.numel() != 95:
88
+ raise ValueError("Expected 95 features")
89
+
90
+ x = x.view(19, 5) # (channels, time)
91
+ x = x.unsqueeze(0) # (batch, 19, 5)
92
+ x = x.to(DEVICE)
93
 
94
  with torch.no_grad():
95
  outputs = model(x)
 
101
  return pred, confidence, probs.tolist()
102
 
103
  # =========================================================
104
+ # ROUTES
105
  # =========================================================
106
 
107
  @app.get("/")
108
  def home():
109
+ return {"message": "NeuroHealth EEG API Running"}
 
 
110
 
111
  # =========================================================
112
+ # ALZHEIMER
113
  # =========================================================
114
 
115
  @app.post("/predict/alzheimer")
 
127
  }
128
 
129
  # =========================================================
130
+ # PARKINSON
131
  # =========================================================
132
 
133
  @app.post("/predict/parkinson")