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 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203
| import torch import torch.nn as nn import torch.nn.functional as F
class InterModalMutualGuidance(nn.Module): """ BiGSTF-Net 跨模态相互引导模块 核心思想: - EEG的时域特征引导fNIRS的空间注意力(EEG告诉fNIRS"看哪里") - fNIRS的空间特征引导EEG的时间注意力(fNIRS告诉EEG"关注什么时段") 这创造了双向信息流,比简单拼接更有效 """ def __init__(self, eeg_channels: int = 32, fnirs_channels: int = 16, hidden_dim: int = 128): super().__init__() self.eeg_to_fnirs_guide = nn.Sequential( nn.Linear(eeg_channels, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, fnirs_channels), nn.Sigmoid() ) self.fnirs_to_eeg_guide = nn.Sequential( nn.Linear(fnirs_channels, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, eeg_channels), nn.Sigmoid() ) self.eeg_proj = nn.Linear(eeg_channels, hidden_dim) self.fnirs_proj = nn.Linear(fnirs_channels, hidden_dim) self.fusion_gate = nn.Sequential( nn.Linear(hidden_dim * 2, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, hidden_dim), nn.Sigmoid() ) def forward(self, eeg_feat: torch.Tensor, fnirs_feat: torch.Tensor) -> tuple: """ Args: eeg_feat: (B, C_eeg, T_eeg) EEG时域特征 fnirs_feat: (B, C_fnirs, T_fnirs) fNIRS时域特征 Returns: guided_eeg: EEG被fNIRS引导后的特征 guided_fnirs: fNIRS被EEG引导后的特征 """ eeg_summary = eeg_feat.mean(dim=-1) fnirs_summary = fnirs_feat.mean(dim=-1) fnirs_spatial_attn = self.eeg_to_fnirs_guide(eeg_summary) guided_fnirs = fnirs_feat * fnirs_spatial_attn.unsqueeze(-1) eeg_channel_attn = self.fnirs_to_eeg_guide(fnirs_summary) guided_eeg = eeg_feat * eeg_channel_attn.unsqueeze(-1) eeg_proj = self.eeg_proj(guided_eeg.mean(dim=-1)) fnirs_proj = self.fnirs_proj(guided_fnirs.mean(dim=-1)) gate = self.fusion_gate(torch.cat([eeg_proj, fnirs_proj], dim=-1)) fused = gate * eeg_proj + (1 - gate) * fnirs_proj return guided_eeg, guided_fnirs, fused
class BiGSTFNet(nn.Module): """ BiGSTF-Net 完整模型 Bi-directional Guided Spatio-Temporal Fusion Network 用于EEG-fNIRS双模态认知状态分类 应用场景: 1. 驾驶员疲劳检测(EEG即时+fNIRS持续) 2. 认知负荷评估 3. 情绪识别 """ def __init__(self, eeg_channels: int = 32, fnirs_channels: int = 16, num_classes: int = 3, hidden_dim: int = 128): super().__init__() self.eeg_encoder = nn.Sequential( nn.Conv1d(eeg_channels, 64, kernel_size=7, stride=2, padding=3), nn.BatchNorm1d(64), nn.ELU(), nn.Conv1d(64, 64, kernel_size=1), nn.BatchNorm1d(64), nn.ELU(), nn.LSTM(64, hidden_dim, batch_first=True, bidirectional=True), ) self.fnirs_encoder = nn.Sequential( nn.Conv1d(fnirs_channels, 32, kernel_size=5, stride=1, padding=2), nn.BatchNorm1d(32), nn.ELU(), nn.Conv1d(32, 32, kernel_size=3, stride=1, padding=1), nn.BatchNorm1d(32), nn.ELU(), ) self.guidance = InterModalMutualGuidance( eeg_channels=hidden_dim * 2, fnirs_channels=32, hidden_dim=hidden_dim ) self.temporal_fusion = nn.TransformerEncoderLayer( d_model=hidden_dim, nhead=8, dim_feedforward=hidden_dim * 2, dropout=0.1, batch_first=True ) self.classifier = nn.Sequential( nn.Linear(hidden_dim, hidden_dim // 2), nn.ELU(), nn.Dropout(0.3), nn.Linear(hidden_dim // 2, num_classes) ) def forward(self, eeg: torch.Tensor, fnirs: torch.Tensor) -> torch.Tensor: """ Args: eeg: (B, C_eeg, T_eeg) EEG信号 fnirs: (B, C_fnirs, T_fnirs) fNIRS信号 Returns: logits: (B, num_classes) 认知状态分类 """ eeg_feat = self.eeg_encoder(eeg) if isinstance(eeg_feat, tuple): eeg_feat, _ = eeg_feat eeg_feat = eeg_feat.transpose(1, 2) fnirs_feat = self.fnirs_encoder(fnirs) guided_eeg, guided_fnirs, fused = self.guidance(eeg_feat, fnirs_feat) fused = fused.unsqueeze(1) fused = self.temporal_fusion(fused) fused = fused.squeeze(1) logits = self.classifier(fused) return logits
if __name__ == "__main__": model = BiGSTFNet( eeg_channels=32, fnirs_channels=16, num_classes=3, hidden_dim=128 ) eeg = torch.randn(4, 32, 256) fnirs = torch.randn(4, 16, 100) logits = model(eeg, fnirs) print(f"输入 EEG shape: {eeg.shape}") print(f"输入 fNIRS shape: {fnirs.shape}") print(f"输出 logits shape: {logits.shape}") print(f"预测: {torch.argmax(logits, dim=-1).tolist()}") print("✅ BiGSTF-Net 测试通过")
|