Spaces:
Sleeping
Sleeping
File size: 5,784 Bytes
403dcd0 b6ccded a0f6a7d 403dcd0 e5f75e4 ebae7d1 e5f75e4 099d98d f4e8a6f b6ccded ebae7d1 403dcd0 72973d5 ebae7d1 b6ccded ebae7d1 b6ccded ebae7d1 b6ccded ebae7d1 b6ccded 72973d5 ebae7d1 72973d5 ebae7d1 b6ccded ebae7d1 b6ccded ebae7d1 925fdc0 ebae7d1 72973d5 b6ccded ebae7d1 b6ccded ebae7d1 72973d5 ebae7d1 72973d5 099d98d 72973d5 099d98d b6ccded 72973d5 b6ccded 72973d5 403dcd0 b6ccded ebae7d1 925fdc0 ebae7d1 925fdc0 ebae7d1 925fdc0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 | import torch
import torch.nn as nn
from torch.nn.utils.parametrizations import weight_norm
class AttentionFusion(nn.Module):
def __init__(self, feature_dim: int, attention_dim: int):
super().__init__()
self.attention = nn.Linear(feature_dim, attention_dim)
self.context_vector = nn.Linear(attention_dim, 1, bias=False)
def forward(self, features: torch.Tensor) -> torch.Tensor:
# features: [batch_size, feature_dim]
attention_scores = self.context_vector(torch.tanh(self.attention(features))) # [batch_size, 1]
attention_weights = torch.softmax(attention_scores, dim=1) # [batch_size, 1]
weighted_features = features * attention_weights # [batch_size, feature_dim]
return weighted_features
class ParkinsonMultimodalEncoder(nn.Module):
def __init__(self,
mri_channels=1,
clinical_dim=8,
datscan_dim=4,
latent_dim=64): # Final latent embedding dimension
super().__init__()
# -------- MRI Pathway (input: [B, 1, D, H, W]) --------
self.mri_encoder = nn.Sequential(
nn.Conv3d(mri_channels, 16, kernel_size=3, padding=1), # [B, 16, D, H, W]
nn.BatchNorm3d(16),
nn.ReLU(),
nn.MaxPool3d(2), # [B, 16, D/2, H/2, W/2]
self._make_res_block3d(16, 32, stride=2), # [B, 32, D/4, H/4, W/4]
self._make_res_block3d(32, 64, stride=1), # [B, 64, D/4, H/4, W/4]
nn.Conv3d(64, 64, kernel_size=1), # [B, 64, D/4, H/4, W/4]
nn.Sigmoid(),
nn.AdaptiveAvgPool3d(1) # [B, 64, 1, 1, 1]
)
# -------- Clinical Pathway (input: [B, T, 8]) --------
self.clinical_norm = nn.LayerNorm(clinical_dim)
self.clinical_conv = weight_norm(nn.Conv1d(8, 16, kernel_size=3, padding=1)) # [B, 16, T]
self.clinical_motor_branch = nn.Sequential(nn.Linear(16, 32), nn.ReLU()) # For time step 0
self.clinical_nonmotor_branch = nn.Sequential(nn.Linear(16, 32), nn.ReLU()) # For time step 1
# -------- DaTscan Pathway (input: [B, 4]) --------
self.datscan_base = nn.Sequential(
nn.Linear(datscan_dim, 16), nn.ReLU(), # [B, 16]
nn.Linear(16, 8), nn.Tanh() # [B, 8]
)
self.putamen_head = nn.Sequential(nn.Linear(8, 4), nn.ReLU()) # [B, 4]
self.caudate_head = nn.Sequential(nn.Linear(8, 4), nn.ReLU()) # [B, 4]
self.ratio_head = nn.Linear(8, 4) # [B, 4]
# -------- Fusion --------
fusion_input_dim = 64 + 64 + 12 # MRI + Clinical + DaTscan = 140
self.attention_fusion = AttentionFusion(feature_dim=fusion_input_dim, attention_dim=64)
self.fusion = nn.Sequential(
nn.Linear(fusion_input_dim, 256), # [B, 256]
nn.LeakyReLU(0.2),
nn.Dropout(0.3),
nn.Linear(256, 64), # [B, 64]
nn.Sigmoid(),
nn.Linear(64, latent_dim) # [B, latent_dim]
)
self._init_weights()
def _make_res_block3d(self, in_channels, out_channels, stride=1):
return nn.Sequential(
nn.Conv3d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1),
nn.BatchNorm3d(out_channels),
nn.ReLU(),
nn.Conv3d(out_channels, out_channels, kernel_size=3, padding=1),
nn.BatchNorm3d(out_channels)
)
def _init_weights(self):
for m in self.modules():
if isinstance(m, nn.Conv3d):
nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
if m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.Linear):
nn.init.xavier_uniform_(m.weight)
if m.bias is not None:
nn.init.constant_(m.bias, 0.01)
# ----- Process each modality individually -----
def process_mri(self, mri: torch.Tensor) -> torch.Tensor:
x = self.mri_encoder(mri) # [B, 64, 1, 1, 1]
return x.view(x.size(0), -1) # [B, 64]
def process_clinical(self, clinical: torch.Tensor) -> torch.Tensor:
clinical = self.clinical_norm(clinical) # [B, T, 8]
clinical = clinical.permute(0, 2, 1) # [B, 8, T]
clinical = self.clinical_conv(clinical) # [B, 16, T]
clinical = clinical.permute(0, 2, 1) # [B, T, 16]
motor = self.clinical_motor_branch(clinical[:, 0, :]) # [B, 32]
nonmotor = self.clinical_nonmotor_branch(clinical[:, 1, :]) # [B, 32]
return torch.cat([motor, nonmotor], dim=1) # [B, 64]
def process_datscan(self, datscan: torch.Tensor) -> torch.Tensor:
base = self.datscan_base(datscan) # [B, 8]
putamen = self.putamen_head(base) # [B, 4]
caudate = self.caudate_head(base) # [B, 4]
ratio = self.ratio_head(base) # [B, 4]
return torch.cat([putamen, caudate, ratio], dim=1) # [B, 12]
def forward(self, mri, clinical, datscan):
# Each returns [B, X]
mri_feat = self.process_mri(mri) # [B, 64]
clinical_feat = self.process_clinical(clinical) # [B, 64]
dat_feat = self.process_datscan(datscan) # [B, 12]
combined = torch.cat([mri_feat, clinical_feat, dat_feat], dim=1) # [B, 140]
weighted = self.attention_fusion(combined) # [B, 140]
latent = self.fusion(weighted) # [B, latent_dim]
return latent |