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
| import torch import torch.nn as nn
class CRNN3D(nn.Module): """ 3D Convolutional Recurrent Neural Network for EEG Decoding 复现论文核心架构: 3D-CRNN - 3D卷积层: 捕获电极空间拓扑 + 时间局部特征 - GRU循环层: 建模时序依赖 - 分类头: 联合预测 RP 和 DI 论文: arXiv:2609.07128 """ def __init__( self, n_channels: int = 62, n_times: int = 500, n_classes_rp: int = 2, n_classes_di: int = 3, dropout: float = 0.5 ): super().__init__() self.conv3d_block = nn.Sequential( nn.Conv3d(1, 32, kernel_size=(7, 7, 5), padding=(3, 3, 2)), nn.BatchNorm3d(32), nn.ELU(), nn.MaxPool3d(kernel_size=(2, 2, 2)), nn.Dropout3d(dropout * 0.5), nn.Conv3d(32, 64, kernel_size=(5, 5, 3), padding=(2, 2, 1)), nn.BatchNorm3d(64), nn.ELU(), nn.MaxPool3d(kernel_size=(2, 2, 2)), nn.Dropout3d(dropout * 0.5), nn.Conv3d(64, 128, kernel_size=(3, 3, 3), padding=(1, 1, 1)), nn.BatchNorm3d(128), nn.ELU(), nn.AdaptiveAvgPool3d((1, 1, None)), ) self.feature_dim = 128 self.gru = nn.GRU( input_size=self.feature_dim, hidden_size=128, num_layers=2, batch_first=True, bidirectional=True, dropout=dropout ) self.rp_head = nn.Sequential( nn.Linear(256, 64), nn.ELU(), nn.Dropout(dropout), nn.Linear(64, n_classes_rp) ) self.di_head = nn.Sequential( nn.Linear(256, 64), nn.ELU(), nn.Dropout(dropout), nn.Linear(64, n_classes_di) ) def forward(self, x: torch.Tensor) -> tuple: """ 前向传播 Args: x: EEG输入, shape=(B, 1, C, T) C=62通道, T=500时间点(1s@500Hz) Returns: rp_out: 风险预测logits, shape=(B, n_classes_rp) di_out: 危险识别logits, shape=(B, n_classes_di) """ x = self.conv3d_block(x) x = x.squeeze(3).squeeze(2) x = x.permute(0, 2, 1) gru_out, _ = self.gru(x) last_hidden = gru_out[:, -1, :] rp_out = self.rp_head(last_hidden) di_out = self.di_head(last_hidden) return rp_out, di_out
if __name__ == "__main__": model = CRNN3D(n_channels=62, n_times=500) n_params = sum(p.numel() for p in model.parameters()) print(f"模型参数量: {n_params:,} ({n_params/1e6:.2f}M)") batch_size = 16 x = torch.randn(batch_size, 1, 62, 500) rp_out, di_out = model(x) print(f"输入: {x.shape}") print(f"RP输出: {rp_out.shape} (二分类)") print(f"DI输出: {di_out.shape} (三分类)") from torch.profiler import profile with profile(activities=[]) as prof: rp_out, di_out = model(x)
|