eyupipler commited on
Commit
f40468f
·
verified ·
1 Parent(s): 9493c86

Upload vbai-2.6ad_test.py

Browse files
Files changed (1) hide show
  1. vbai-2.6ad_test.py +413 -0
vbai-2.6ad_test.py ADDED
@@ -0,0 +1,413 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Vbai-2.6AD Standalone Inference Script
3
+ ===============================================
4
+ A self-contained script containing the full model architecture and inference logic.
5
+ Designed for open-source distribution.
6
+
7
+ Usage:
8
+ python vbai-2.6ad_test.py --weights path/to/model.pt --mri path/to/scan.nii --clinical "Age:75.2, Sex:1, MMSE:25"
9
+ """
10
+ import argparse
11
+ import os
12
+ import torch
13
+ import torch.nn as nn
14
+ import torch.nn.functional as F
15
+ import numpy as np
16
+
17
+ try:
18
+ import nibabel as nib
19
+ from scipy.ndimage import zoom
20
+ HAS_NIBABEL = True
21
+ except ImportError:
22
+ HAS_NIBABEL = False
23
+
24
+ # ============================================================
25
+ # Configuration Constants
26
+ # ============================================================
27
+ FEATURE_NAMES = [
28
+ "Age", "Sex", "MMSE", "CDRSB", "APOE4_count",
29
+ "CSF_ABETA42", "CSF_TAU", "CSF_PTAU", "CSF_AB42_AB40",
30
+ "PLASMA_PTAU", "PLASMA_NFL", "PLASMA_AB42_AB40", "PLASMA_GFAP"
31
+ ]
32
+ CLASS_NAMES = ["CN", "MCI", "AD"]
33
+
34
+ class ModelConfig:
35
+ def __init__(self):
36
+ self.mri_input_shape = (1, 96, 96, 96)
37
+ self.mri_encoder_channels = [32, 64, 128, 256]
38
+ self.mri_bottleneck_channels = 512
39
+ self.mri_feature_dim = 512
40
+ self.mri_dropout = 0.4
41
+ self.use_cbam = True
42
+ self.use_se_block = True
43
+
44
+ self.num_tabular_inputs = len(FEATURE_NAMES) * 2
45
+ self.tabular_hidden_dims = [128, 256]
46
+ self.tabular_feature_dim = 256
47
+ self.tabular_dropout = 0.3
48
+
49
+ self.fusion_dim = 512
50
+ self.fusion_num_heads = 8
51
+ self.fusion_dropout = 0.3
52
+
53
+ self.num_classes = 3
54
+ self.progression_hidden_dim = 256
55
+ self.max_progression_months = 120
56
+ self.num_time_bins = 24
57
+
58
+ # ============================================================
59
+ # Attention Modules
60
+ # ============================================================
61
+ class ChannelAttention3D(nn.Module):
62
+ def __init__(self, ch, r=16):
63
+ super().__init__()
64
+ m = max(ch // r, 8)
65
+ self.mlp = nn.Sequential(nn.Linear(ch, m), nn.ReLU(inplace=True), nn.Linear(m, ch))
66
+
67
+ def forward(self, x):
68
+ a = x.mean(dim=[2, 3, 4]); b = x.amax(dim=[2, 3, 4])
69
+ attn = torch.sigmoid(self.mlp(a) + self.mlp(b))
70
+ return x * attn[..., None, None, None]
71
+
72
+ class SpatialAttention3D(nn.Module):
73
+ def __init__(self, k=7):
74
+ super().__init__()
75
+ self.conv = nn.Conv3d(2, 1, k, padding=k // 2, bias=False)
76
+
77
+ def forward(self, x):
78
+ avg = x.mean(dim=1, keepdim=True); mx = x.amax(dim=1, keepdim=True)
79
+ attn = torch.sigmoid(self.conv(torch.cat([avg, mx], dim=1)))
80
+ return x * attn
81
+
82
+ class CBAM3D(nn.Module):
83
+ def __init__(self, ch, r=16):
84
+ super().__init__()
85
+ self.c = ChannelAttention3D(ch, r); self.s = SpatialAttention3D()
86
+ def forward(self, x): return self.s(self.c(x))
87
+
88
+ class SEBlock3D(nn.Module):
89
+ def __init__(self, ch, r=16):
90
+ super().__init__()
91
+ m = max(ch // r, 8)
92
+ self.fc = nn.Sequential(nn.Linear(ch, m), nn.ReLU(True), nn.Linear(m, ch), nn.Sigmoid())
93
+ def forward(self, x):
94
+ s = x.mean(dim=[2, 3, 4]); s = self.fc(s)[..., None, None, None]
95
+ return x * s
96
+
97
+ # ============================================================
98
+ # Encoders
99
+ # ============================================================
100
+ class ResBlock3D(nn.Module):
101
+ def __init__(self, in_ch, out_ch, stride=1, use_cbam=True, use_se=True, drop_path=0.0):
102
+ super().__init__()
103
+ self.conv1 = nn.Conv3d(in_ch, out_ch, 3, stride, 1, bias=False)
104
+ self.bn1 = nn.BatchNorm3d(out_ch)
105
+ self.conv2 = nn.Conv3d(out_ch, out_ch, 3, 1, 1, bias=False)
106
+ self.bn2 = nn.BatchNorm3d(out_ch)
107
+ self.act = nn.GELU()
108
+ self.cbam = CBAM3D(out_ch) if use_cbam else nn.Identity()
109
+ self.se = SEBlock3D(out_ch) if use_se else nn.Identity()
110
+ self.skip = nn.Identity() if (in_ch == out_ch and stride == 1) else nn.Sequential(
111
+ nn.Conv3d(in_ch, out_ch, 1, stride, bias=False), nn.BatchNorm3d(out_ch))
112
+
113
+ def forward(self, x):
114
+ identity = self.skip(x)
115
+ out = self.act(self.bn1(self.conv1(x)))
116
+ out = self.bn2(self.conv2(out))
117
+ out = self.cbam(out); out = self.se(out)
118
+ return self.act(out + identity)
119
+
120
+ class ASPP3D(nn.Module):
121
+ def __init__(self, in_ch, out_ch, dilations=(1, 6, 12, 18)):
122
+ super().__init__()
123
+ per = out_ch // len(dilations)
124
+ self.branches = nn.ModuleList([
125
+ nn.Sequential(nn.Conv3d(in_ch, per, 3, padding=d, dilation=d, bias=False),
126
+ nn.BatchNorm3d(per), nn.GELU())
127
+ for d in dilations
128
+ ])
129
+ self.gp = nn.Sequential(
130
+ nn.AdaptiveAvgPool3d(1),
131
+ nn.Conv3d(in_ch, per, 1, bias=False),
132
+ nn.BatchNorm3d(per), nn.GELU())
133
+ self.fuse = nn.Sequential(nn.Conv3d(per * (len(dilations) + 1), out_ch, 1, bias=False),
134
+ nn.BatchNorm3d(out_ch), nn.GELU())
135
+
136
+ def forward(self, x):
137
+ feats = [b(x) for b in self.branches]
138
+ g = self.gp(x)
139
+ g = F.interpolate(g, size=x.shape[2:], mode="trilinear", align_corners=False)
140
+ feats.append(g)
141
+ return self.fuse(torch.cat(feats, dim=1))
142
+
143
+ class MRIEncoder3D(nn.Module):
144
+ def __init__(self, mcfg: ModelConfig):
145
+ super().__init__()
146
+ ch = mcfg.mri_encoder_channels
147
+ self.stem = nn.Sequential(
148
+ nn.Conv3d(1, ch[0], 7, 2, 3, bias=False), nn.BatchNorm3d(ch[0]), nn.GELU(),
149
+ nn.MaxPool3d(3, 2, 1))
150
+ self.stage1 = self._make(ch[0], ch[0], 2, 1, mcfg)
151
+ self.stage2 = self._make(ch[0], ch[1], 2, 2, mcfg)
152
+ self.stage3 = self._make(ch[1], ch[2], 2, 2, mcfg)
153
+ self.stage4 = self._make(ch[2], ch[3], 2, 2, mcfg)
154
+ self.aspp = ASPP3D(ch[3], mcfg.mri_bottleneck_channels)
155
+ self.pool = nn.AdaptiveAvgPool3d(1)
156
+ self.proj = nn.Sequential(
157
+ nn.Linear(mcfg.mri_bottleneck_channels, mcfg.mri_feature_dim),
158
+ nn.GELU(), nn.Dropout(mcfg.mri_dropout))
159
+
160
+ def _make(self, in_ch, out_ch, n, stride, mcfg):
161
+ layers = [ResBlock3D(in_ch, out_ch, stride, mcfg.use_cbam, mcfg.use_se_block)]
162
+ for _ in range(1, n):
163
+ layers.append(ResBlock3D(out_ch, out_ch, 1, mcfg.use_cbam, mcfg.use_se_block))
164
+ return nn.Sequential(*layers)
165
+
166
+ def forward(self, x):
167
+ x = self.stem(x)
168
+ x = self.stage1(x); x = self.stage2(x); x = self.stage3(x); x = self.stage4(x)
169
+ x = self.aspp(x); x = self.pool(x).flatten(1)
170
+ return self.proj(x)
171
+
172
+ class TabularEncoder(nn.Module):
173
+ def __init__(self, mcfg: ModelConfig):
174
+ super().__init__()
175
+ prev = mcfg.num_tabular_inputs
176
+ layers = []
177
+ for h in mcfg.tabular_hidden_dims:
178
+ layers += [nn.Linear(prev, h), nn.LayerNorm(h), nn.GELU(), nn.Dropout(mcfg.tabular_dropout)]
179
+ prev = h
180
+ layers += [nn.Linear(prev, mcfg.tabular_feature_dim)]
181
+ self.net = nn.Sequential(*layers)
182
+
183
+ def forward(self, x):
184
+ return self.net(x)
185
+
186
+ # ============================================================
187
+ # Fusion & Heads
188
+ # ============================================================
189
+ class CrossModalFusion(nn.Module):
190
+ def __init__(self, mri_dim, tab_dim, fdim, heads=8, dropout=0.1):
191
+ super().__init__()
192
+ self.pm = nn.Linear(mri_dim, fdim); self.pt = nn.Linear(tab_dim, fdim)
193
+ self.a_mt = nn.MultiheadAttention(fdim, heads, dropout=dropout, batch_first=True)
194
+ self.a_tm = nn.MultiheadAttention(fdim, heads, dropout=dropout, batch_first=True)
195
+ self.lnm = nn.LayerNorm(fdim); self.lnt = nn.LayerNorm(fdim)
196
+ self.gate = nn.Sequential(nn.Linear(fdim * 2, fdim), nn.Sigmoid())
197
+ self.out = nn.Sequential(nn.Linear(fdim * 2, fdim), nn.GELU(), nn.Dropout(dropout))
198
+
199
+ def forward(self, m, t):
200
+ m1 = self.pm(m).unsqueeze(1); t1 = self.pt(t).unsqueeze(1)
201
+ ma, _ = self.a_mt(m1, t1, t1); ta, _ = self.a_tm(t1, m1, m1)
202
+ m2 = self.lnm(m1 + ma).squeeze(1); t2 = self.lnt(t1 + ta).squeeze(1)
203
+ cat = torch.cat([m2, t2], dim=-1)
204
+ g = self.gate(cat); o = self.out(cat)
205
+ return g * m2 + (1 - g) * t2 + o
206
+
207
+ class ClsHead(nn.Module):
208
+ def __init__(self, in_dim, num_classes, dropout=0.3):
209
+ super().__init__()
210
+ self.h = nn.Sequential(
211
+ nn.Linear(in_dim, 256), nn.GELU(), nn.Dropout(dropout),
212
+ nn.Linear(256, 128), nn.GELU(), nn.Dropout(dropout),
213
+ nn.Linear(128, num_classes))
214
+ def forward(self, x): return self.h(x)
215
+
216
+ class ProgressionHead(nn.Module):
217
+ def __init__(self, in_dim, hidden=256, max_months=120, n_bins=24):
218
+ super().__init__()
219
+ self.max_months = float(max_months); self.n_bins = n_bins
220
+ self.shared = nn.Sequential(nn.Linear(in_dim, hidden), nn.GELU(), nn.Dropout(0.3))
221
+ self.binary = nn.Linear(hidden, 1)
222
+ self.time = nn.Sequential(nn.Linear(hidden, 64), nn.GELU(), nn.Linear(64, 1))
223
+
224
+ def forward(self, x):
225
+ h = self.shared(x)
226
+ logits = self.binary(h).squeeze(-1)
227
+ return {
228
+ "will_progress": torch.sigmoid(logits),
229
+ "time_to_conversion": torch.clamp(F.softplus(self.time(h)).squeeze(-1), min=0.0, max=self.max_months),
230
+ }
231
+
232
+ # ============================================================
233
+ # Main Model Class
234
+ # ============================================================
235
+ class HFv3AModel(nn.Module):
236
+ def __init__(self, mcfg: ModelConfig = None):
237
+ super().__init__()
238
+ self.cfg = mcfg or ModelConfig()
239
+ self.mri_encoder = MRIEncoder3D(self.cfg)
240
+ self.tab_encoder = TabularEncoder(self.cfg)
241
+ self.mri_classifier = ClsHead(self.cfg.mri_feature_dim, self.cfg.num_classes, self.cfg.mri_dropout)
242
+ self.tab_classifier = ClsHead(self.cfg.tabular_feature_dim, self.cfg.num_classes, self.cfg.tabular_dropout)
243
+ self.fusion = CrossModalFusion(self.cfg.mri_feature_dim, self.cfg.tabular_feature_dim, self.cfg.fusion_dim, self.cfg.fusion_num_heads, self.cfg.fusion_dropout)
244
+ self.fused_classifier = ClsHead(self.cfg.fusion_dim, self.cfg.num_classes, self.cfg.fusion_dropout)
245
+ self.progression_head = ProgressionHead(self.cfg.fusion_dim, self.cfg.progression_hidden_dim, self.cfg.max_progression_months, self.cfg.num_time_bins)
246
+
247
+ def forward(self, mri=None, tab=None):
248
+ out = {}
249
+ m_feat = t_feat = None
250
+ if mri is not None:
251
+ m_feat = self.mri_encoder(mri)
252
+ out["mri_logits"] = self.mri_classifier(m_feat)
253
+ if tab is not None:
254
+ t_feat = self.tab_encoder(tab)
255
+ out["tab_logits"] = self.tab_classifier(t_feat)
256
+
257
+ if m_feat is not None and t_feat is not None:
258
+ f = self.fusion(m_feat, t_feat)
259
+ out["fused_logits"] = self.fused_classifier(f)
260
+ out["progression"] = self.progression_head(f)
261
+ elif m_feat is not None:
262
+ out["fused_logits"] = out["mri_logits"]
263
+ elif t_feat is not None:
264
+ out["fused_logits"] = out["tab_logits"]
265
+
266
+ return out
267
+
268
+ @torch.no_grad()
269
+ def predict(self, mri=None, tab=None):
270
+ self.eval()
271
+ out = self.forward(mri=mri, tab=tab)
272
+ probs = F.softmax(out["fused_logits"], dim=-1)
273
+ pred = probs.argmax(dim=-1)
274
+ result = {
275
+ "pred_class": pred,
276
+ "class_probs": probs,
277
+ "class_name": CLASS_NAMES[pred.item()]
278
+ }
279
+ if "progression" in out:
280
+ p = out["progression"]
281
+ result["will_progress"] = p["will_progress"].item()
282
+ result["time_to_conversion_months"] = p["time_to_conversion"].item()
283
+ return result
284
+
285
+ # ============================================================
286
+ # Preprocessing Helpers
287
+ # ============================================================
288
+ def load_mri_tensor(path: str, target_shape=(96, 96, 96)):
289
+ if not HAS_NIBABEL:
290
+ raise ImportError("Please install nibabel and scipy to process NIfTI MRI images: pip install nibabel scipy")
291
+
292
+ img = nib.load(path)
293
+ data = img.get_fdata().astype(np.float32)
294
+ if data.ndim == 4:
295
+ data = data[..., 0]
296
+
297
+ mask = data > 0
298
+ if mask.sum() > 0:
299
+ vals = data[mask]
300
+ lo, hi = np.percentile(vals, [1.0, 99.0])
301
+ data = np.clip(data, lo, hi)
302
+ m, s = vals.mean(), vals.std()
303
+ if s > 0:
304
+ data = (data - m) / s
305
+ data[~mask] = 0
306
+
307
+ if data.shape != target_shape:
308
+ f = [t / s for t, s in zip(target_shape, data.shape)]
309
+ data = zoom(data, f, order=1)
310
+
311
+ tensor = torch.from_numpy(np.ascontiguousarray(data)).unsqueeze(0).unsqueeze(0).float()
312
+ return tensor
313
+
314
+ def parse_clinical_data(clinical_str: str):
315
+ """
316
+ Parses a string like "Age:75.2, Sex:1, MMSE:25" into a tabular tensor.
317
+ If normalizer values are not provided, it passes raw values.
318
+ """
319
+ pairs = [p.strip().split(':') for p in clinical_str.split(',') if ':' in p]
320
+ val_dict = {k.strip(): float(v.strip()) for k, v in pairs}
321
+
322
+ vals = []
323
+ masks = []
324
+ for feat in FEATURE_NAMES:
325
+ if feat in val_dict:
326
+ vals.append(val_dict[feat])
327
+ masks.append(1.0)
328
+ else:
329
+ vals.append(0.0)
330
+ masks.append(0.0)
331
+
332
+ # Note: Without the original training normalizer state, scaling might be inaccurate.
333
+ # We pass the raw values here. For accurate production use, you should apply your normalizer parameters.
334
+ tab_tensor = torch.tensor(vals + masks, dtype=torch.float32).unsqueeze(0)
335
+ return tab_tensor
336
+
337
+ # ============================================================
338
+ # Inference Pipeline
339
+ # ============================================================
340
+ def run_inference(args):
341
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
342
+ print(f"[*] Running on device: {device}")
343
+
344
+ print("[*] Initializing Model...")
345
+ model = HFv3AModel().to(device)
346
+ total_params = sum(p.numel() for p in model.parameters())
347
+ print(f"[*] Model Parameters: {total_params:,} ({(total_params/1e6):.2f}M)")
348
+
349
+ if args.weights:
350
+ if os.path.exists(args.weights):
351
+ print(f"[*] Loading weights from {args.weights}")
352
+ try:
353
+ ckpt = torch.load(args.weights, map_location=device, weights_only=False)
354
+ # Flexible loading depending on how weights were saved
355
+ state_dict = ckpt["model"] if "model" in ckpt else (ckpt["model_state_dict"] if "model_state_dict" in ckpt else ckpt)
356
+ model.load_state_dict(state_dict, strict=False)
357
+ except Exception as e:
358
+ print(f"[!] Error loading weights: {e}")
359
+ else:
360
+ print(f"[!] Warning: Weights file not found at {args.weights}. Using untrained model.")
361
+
362
+ model.eval()
363
+
364
+ mri_tensor = None
365
+ tab_tensor = None
366
+
367
+ if args.mri:
368
+ if os.path.exists(args.mri):
369
+ print(f"[*] Processing MRI: {args.mri}")
370
+ mri_tensor = load_mri_tensor(args.mri).to(device)
371
+ else:
372
+ print(f"[!] Error: MRI file not found at {args.mri}")
373
+ return
374
+
375
+ if args.clinical:
376
+ print(f"[*] Processing Clinical Data: {args.clinical}")
377
+ tab_tensor = parse_clinical_data(args.clinical).to(device)
378
+
379
+ if mri_tensor is None and tab_tensor is None:
380
+ print("[!] Error: You must provide either --mri or --clinical inputs.")
381
+ return
382
+
383
+ print("[*] Running Prediction...")
384
+ result = model.predict(mri=mri_tensor, tab=tab_tensor)
385
+
386
+ print("\n" + "="*40)
387
+ print(" PREDICTION RESULTS ")
388
+ print("="*40)
389
+ print(f"Diagnosis : {result['class_name']} (Class {result['pred_class'].item()})")
390
+
391
+ probs = result['class_probs'].squeeze().tolist()
392
+ print(f"Confidence (CN) : {probs[0]:.4f}")
393
+ print(f"Confidence (MCI) : {probs[1]:.4f}")
394
+ print(f"Confidence (AD) : {probs[2]:.4f}")
395
+
396
+ if "will_progress" in result:
397
+ print("-" * 40)
398
+ print(f"Progression Risk : {result['will_progress']:.2%}")
399
+ print(f"Est. Time to Convert: {result['time_to_conversion_months']:.1f} months")
400
+ print("="*40)
401
+
402
+ if __name__ == "__main__":
403
+ parser = argparse.ArgumentParser(description="Vbai-2.6AD Inference Script")
404
+ parser.add_argument("--weights", type=str, help="Path to the model .pt weights file")
405
+ parser.add_argument("--mri", type=str, help="Path to the input NIfTI (.nii / .nii.gz) MRI scan")
406
+ parser.add_argument("--clinical", type=str, help='Clinical data string, e.g., "Age:75.2, Sex:1, MMSE:25"')
407
+
408
+ args = parser.parse_args()
409
+
410
+ if not any([args.mri, args.clinical]):
411
+ parser.print_help()
412
+ else:
413
+ run_inference(args)