Mikecode123 commited on
Commit
86290ef
·
verified ·
1 Parent(s): 190f3ca

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +45 -39
app.py CHANGED
@@ -10,8 +10,8 @@ import numpy as np
10
  # =========================================================
11
 
12
  app = FastAPI(
13
- title="NeuroHealth EEG CNN API",
14
- version="5.0"
15
  )
16
 
17
  DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
@@ -24,35 +24,33 @@ AD_CLASSES = ["Alzheimer", "FTD", "Control"]
24
  PD_CLASSES = ["Parkinson", "Control"]
25
 
26
  # =========================================================
27
- # REQUEST SCHEMA
28
  # =========================================================
29
 
30
  class EEGRequest(BaseModel):
31
  features: list
32
 
33
  # =========================================================
34
- # CNN MODEL (MATCH YOUR CHECKPOINT)
35
  # =========================================================
36
 
37
- class EEG_CNN(nn.Module):
38
  def __init__(self, output_dim):
39
  super().__init__()
40
 
41
- self.conv1 = nn.Conv1d(19, 32, kernel_size=7, padding=3)
42
  self.bn1 = nn.BatchNorm1d(32)
43
 
44
- self.conv2 = nn.Conv1d(32, 64, kernel_size=5, padding=2)
45
  self.bn2 = nn.BatchNorm1d(64)
46
 
47
- self.conv3 = nn.Conv1d(64, 128, kernel_size=3, padding=1)
48
  self.bn3 = nn.BatchNorm1d(128)
49
 
50
  self.pool = nn.AdaptiveAvgPool1d(1)
51
-
52
  self.fc = nn.Linear(128, output_dim)
53
 
54
  def forward(self, x):
55
- # EXPECTED INPUT: (batch, 19, 76)
56
  x = x.view(x.size(0), 19, -1)
57
 
58
  x = torch.relu(self.bn1(self.conv1(x)))
@@ -60,54 +58,69 @@ class EEG_CNN(nn.Module):
60
  x = torch.relu(self.bn3(self.conv3(x)))
61
 
62
  x = self.pool(x).squeeze(-1)
63
-
64
  return self.fc(x)
65
 
66
  # =========================================================
67
- # CONFIG
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
68
  # =========================================================
69
 
70
  AD_MODEL_PATH = "AD_eeg_cnn_ad_ftd_cn.pt"
71
  PD_MODEL_PATH = "PD_eeg_cnn.pth"
72
 
 
 
73
  # =========================================================
74
  # LOAD MODELS
75
  # =========================================================
76
 
77
- print("Loading AD CNN model...")
78
-
79
- ad_model = EEG_CNN(3).to(DEVICE)
80
  ad_model.load_state_dict(torch.load(AD_MODEL_PATH, map_location=DEVICE))
81
  ad_model.eval()
 
82
 
83
- print("AD model loaded")
84
-
85
- print("Loading PD CNN model...")
86
-
87
- pd_model = EEG_CNN(2).to(DEVICE)
88
  pd_model.load_state_dict(torch.load(PD_MODEL_PATH, map_location=DEVICE))
89
  pd_model.eval()
90
-
91
- print("PD model loaded")
92
-
93
- print("All models ready")
94
 
95
  # =========================================================
96
- # PREDICTION FUNCTION
97
  # =========================================================
98
 
99
  def predict(model, features, classes):
100
 
101
  x = torch.tensor(features, dtype=torch.float32)
102
 
103
- if x.numel() != 19 * 76:
104
- raise ValueError("Expected 19x76 = 1444 features")
105
-
106
- x = x.view(1, 19, 76).to(DEVICE)
107
 
108
  with torch.no_grad():
109
- logits = model(x)
110
- probs = torch.softmax(logits, dim=1).cpu().numpy()[0]
111
 
112
  pred = int(np.argmax(probs))
113
 
@@ -125,14 +138,7 @@ def predict(model, features, classes):
125
 
126
  @app.get("/")
127
  def home():
128
- return {
129
- "status": "NeuroHealth EEG CNN API running",
130
- "input_shape": "19 x 76"
131
- }
132
-
133
- @app.get("/health")
134
- def health():
135
- return {"status": "ok"}
136
 
137
  @app.post("/predict/ad")
138
  def predict_ad(req: EEGRequest):
 
10
  # =========================================================
11
 
12
  app = FastAPI(
13
+ title="NeuroHealth EEG API",
14
+ version="6.0"
15
  )
16
 
17
  DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
 
24
  PD_CLASSES = ["Parkinson", "Control"]
25
 
26
  # =========================================================
27
+ # INPUT
28
  # =========================================================
29
 
30
  class EEGRequest(BaseModel):
31
  features: list
32
 
33
  # =========================================================
34
+ # AD MODEL (conv style)
35
  # =========================================================
36
 
37
+ class EEG_CNN_AD(nn.Module):
38
  def __init__(self, output_dim):
39
  super().__init__()
40
 
41
+ self.conv1 = nn.Conv1d(19, 32, 7, padding=3)
42
  self.bn1 = nn.BatchNorm1d(32)
43
 
44
+ self.conv2 = nn.Conv1d(32, 64, 5, padding=2)
45
  self.bn2 = nn.BatchNorm1d(64)
46
 
47
+ self.conv3 = nn.Conv1d(64, 128, 3, padding=1)
48
  self.bn3 = nn.BatchNorm1d(128)
49
 
50
  self.pool = nn.AdaptiveAvgPool1d(1)
 
51
  self.fc = nn.Linear(128, output_dim)
52
 
53
  def forward(self, x):
 
54
  x = x.view(x.size(0), 19, -1)
55
 
56
  x = torch.relu(self.bn1(self.conv1(x)))
 
58
  x = torch.relu(self.bn3(self.conv3(x)))
59
 
60
  x = self.pool(x).squeeze(-1)
 
61
  return self.fc(x)
62
 
63
  # =========================================================
64
+ # PD MODEL (SEQUENTIAL net style - YOUR CHECKPOINT)
65
+ # =========================================================
66
+
67
+ class EEG_CNN_PD(nn.Module):
68
+ def __init__(self, input_dim, output_dim):
69
+ super().__init__()
70
+
71
+ self.net = nn.Sequential(
72
+ nn.Linear(input_dim, 256), # net.0
73
+ nn.ReLU(),
74
+ nn.BatchNorm1d(256), # net.2
75
+
76
+ nn.Linear(256, 128), # net.4
77
+ nn.ReLU(),
78
+ nn.BatchNorm1d(128), # net.6
79
+
80
+ nn.Linear(128, output_dim) # net.7
81
+ )
82
+
83
+ def forward(self, x):
84
+ return self.net(x)
85
+
86
+ # =========================================================
87
+ # PATHS
88
  # =========================================================
89
 
90
  AD_MODEL_PATH = "AD_eeg_cnn_ad_ftd_cn.pt"
91
  PD_MODEL_PATH = "PD_eeg_cnn.pth"
92
 
93
+ INPUT_DIM = 19 * 76
94
+
95
  # =========================================================
96
  # LOAD MODELS
97
  # =========================================================
98
 
99
+ print("Loading AD model...")
100
+ ad_model = EEG_CNN_AD(3).to(DEVICE)
 
101
  ad_model.load_state_dict(torch.load(AD_MODEL_PATH, map_location=DEVICE))
102
  ad_model.eval()
103
+ print("AD loaded")
104
 
105
+ print("Loading PD model...")
106
+ pd_model = EEG_CNN_PD(INPUT_DIM, 2).to(DEVICE)
 
 
 
107
  pd_model.load_state_dict(torch.load(PD_MODEL_PATH, map_location=DEVICE))
108
  pd_model.eval()
109
+ print("PD loaded")
 
 
 
110
 
111
  # =========================================================
112
+ # PREDICT
113
  # =========================================================
114
 
115
  def predict(model, features, classes):
116
 
117
  x = torch.tensor(features, dtype=torch.float32)
118
 
119
+ x = x.view(1, -1).to(DEVICE)
 
 
 
120
 
121
  with torch.no_grad():
122
+ out = model(x)
123
+ probs = torch.softmax(out, dim=1).cpu().numpy()[0]
124
 
125
  pred = int(np.argmax(probs))
126
 
 
138
 
139
  @app.get("/")
140
  def home():
141
+ return {"status": "running"}
 
 
 
 
 
 
 
142
 
143
  @app.post("/predict/ad")
144
  def predict_ad(req: EEGRequest):