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
| import torch import torch.nn as nn
class Conv3DBlock(nn.Module): """3D卷积块""" def __init__(self, in_channels, out_channels, kernel_size=(3, 3, 3), pool=(2, 2, 2)): super().__init__() self.conv = nn.Conv3d(in_channels, out_channels, kernel_size, padding=kernel_size[0]//2) self.bn = nn.BatchNorm3d(out_channels) self.relu = nn.ReLU() self.pool = nn.MaxPool3d(pool) if pool else nn.Identity() def forward(self, x): return self.pool(self.relu(self.bn(self.conv(x))))
class CRNNModel(nn.Module): """ 3D-CRNN: 3D卷积+循环神经网络 输入: EEG信号 (B, C_eeg, T, H, W) - C_eeg: EEG通道数(如32通道) - T: 时间窗口(如500ms) - H, W: 时空拓扑(如9x9电极位置矩阵) 输出: - RP: 风险预测 二分类 - DI: 危险识别 多分类 """ def __init__(self, eeg_channels=32, num_electrodes=9, hidden_dim=128, num_classes_di=4): super().__init__() self.conv_blocks = nn.ModuleList([ Conv3DBlock(1, 16, (3,3,3), (2,2,1)), Conv3DBlock(16, 32, (3,3,3), (2,2,1)), Conv3DBlock(32, 64, (3,3,3), (2,1,1)), ]) self.flat_dim = 64 * (num_electrodes // 4) * (num_electrodes // 4) self.lstm = nn.LSTM( input_size=self.flat_dim, hidden_size=hidden_dim, num_layers=2, batch_first=True, bidirectional=True ) self.rp_head = nn.Sequential( nn.Linear(hidden_dim * 2, hidden_dim), nn.ReLU(), nn.Dropout(0.3), nn.Linear(hidden_dim, 2) ) self.di_head = nn.Sequential( nn.Linear(hidden_dim * 2, hidden_dim), nn.ReLU(), nn.Dropout(0.3), nn.Linear(hidden_dim, num_classes_di) ) def forward(self, x): """ Args: x: (B, 1, C_eeg, T, H, W) EEG信号 实际输入: (B, 1, 32, 500, 9, 9) for 32ch, 500ms, 9x9 topo Returns: rp_logits: (B, 2) 风险预测 di_logits: (B, num_classes_di) 危险识别 """ for block in self.conv_blocks: x = block(x) B, C, T, H, W = x.shape x = x.permute(0, 2, 1, 3, 4) x = x.reshape(B, T, -1) lstm_out, _ = self.lstm(x) final_state = lstm_out[:, -1, :] rp_logits = self.rp_head(final_state) di_logits = self.di_head(final_state) return rp_logits, di_logits
class RSLLabeler: """ Risk-aware Sequential Labeling 将连续EEG信号标注为序列: - 前N秒标记为"风险预测" - 事件发生时标记为"危险识别" - 中间过渡区标记为"过渡" RSL通过引入软标签平滑提升DI准确率 """ def __init__(self, risk_window=2.0, event_duration=1.0): self.risk_window = risk_window self.event_duration = event_duration def label_sequence(self, events, total_duration, sample_rate=250): """ 生成序列标签 Args: events: [(start_time, end_time, event_type), ...] total_duration: 总时长(秒) sample_rate: EEG采样率 Returns: labels: (N,) 序列标签 0=safe, 1=risk, 2=danger soft_labels: (N, 3) 软标签 """ n_samples = int(total_duration * sample_rate) labels = torch.zeros(n_samples, dtype=torch.long) soft_labels = torch.zeros(n_samples, 3) for start, end, etype in events: s = int(start * sample_rate) e = int(end * sample_rate) risk_s = int((start - self.risk_window) * sample_rate) risk_s = max(0, risk_s) labels[risk_s:s] = 1 soft_labels[risk_s:s, 1] = 0.7 soft_labels[risk_s:s, 0] = 0.3 labels[s:e] = 2 soft_labels[s:e, 2] = 0.9 soft_labels[s:e, 1] = 0.1 safe_mask = labels == 0 soft_labels[safe_mask, 0] = 0.9 soft_labels[safe_mask, 1] = 0.1 return labels, soft_labels
|