
Stanford CS224W: Machine Learning with Graphs
代码下载:https://t.zsxq.com/LhuKn
请索引第43个项目
![]() | ![]() |

数据/任务的动机和解释
脑电图(EEG)是捕捉癫痫发作活动的主要方法,它通过放置在头皮上的电极记录异常脑电活动。了解癫痫发作的特征具有重要的临床意义,因为癫痫发作可能导致短暂的注意力丧失或全身抽搐,而频繁发作会增加患者受伤的风险,甚至可能导致死亡。
以往的研究主要集中于使用经典机器学习方法预测癫痫发作起始时间——一个二元分类任务。然而,预测特定癫痫发作的持续时间仍然是一个相对未被探索的领域,但它具有重要的临床应用价值,从早期检测癫痫持续状态风险到优化神经刺激治疗的剂量。已有研究表明,癫痫发作起始时间可以比随机猜测更准确地预测持续时间,尽管这通常是通过将癫痫发作分为低持续时间和高持续时间两类来实现的(Liu, Y. et al)。本项目旨在利用图神经网络(GNN)和虚拟节点图增强技术来改进这些基线方法,从而更精确地预测特定的持续时间指标,弥合检测和可操作的临床预测之间的差距。将脑电图(EEG)数据表示为图,可以让我们保留大脑电活动的固有拓扑结构,将传感器视为相互连接的节点,而不是孤立的特征。这种结构使 GNN 能够有效地学习驱动癫痫发作演变的复杂空间依赖性和功能连接模式,从而提供比传统时间序列或基于图像的方法更丰富的表示。
数据
脑电图包含来自放置在头皮上的电极的多通道脑活动数据——我们可以根据电极在三维空间中的位置提取空间信息,并从每个节点的连续记录中提取时间信息。节点之间的连接可以通过头皮上的物理位置或节点记录之间的相关性来建模,然后可以利用这些相关性构建记录的图表示。
本研究使用了CHB-MIT脑电图数据集。这是一个公开数据集(https://physionet.org/content/chbmit/1.0.0/),包含22名难治性癫痫患者的脑电图记录,共记录了664次脑电图,记录了198次癫痫发作。脑电图的通道数、电极位置、时间分辨率(采样率)以及电压计算方式各不相同。脑电图电压始终表示为两点之间的差值,因此它可以采用单一参考点作为所有通道的参考,也可以每个通道表示两个相邻电极之间的差值(双极)。我们使用的数据集为双极格式,采样率为每秒256个样本,分辨率为16位,并采用国际10-20脑电图电极位置和命名系统。数据集中的一些样本缺少某些通道或包含一些额外的通道,因此我们仅使用了所有样本中都存在的通道。

CHB-MIT数据集中的癫痫发作探索性数据分析
(我们最初计划使用多个脑电图数据集,但由于记录方法和数据结构的差异,这带来了太大的挑战)。
特征工程和预处理
原始脑电图记录包含来自多个相关通道的复杂波形,代表大脑内的电活动和波动。为了从中提取可推广至多个患者的有用信息,我们执行了一些预处理步骤和特征提取,这些步骤和特征提取在之前的研究中已被证明是有效的(Kerr等人)。对于模型的输入数据,我们配置了发作前持续时间、发作后持续时间以及用于应对临床医生标注错误的缓冲区等参数。更长的发作前和发作后持续时间可能会带来更准确的预测,但会因使用更多数据而增加模型和预测时间。此外,过长的发作前窗口可能会导致早期癫痫发作的数据泄露,而过长的发作后窗口可能会因延迟预测而导致临床应用无用。我们选择10秒的发作前和发作后持续时间,并从数据集中过滤掉所有短于10秒的癫痫发作。
首先,对原始信号应用陷波滤波器以去除工频干扰。由于 CHB-MIT 数据集包含来自美国的信号,因此应用带宽为 6Hz、频率为 60Hz 的滤波器来去除工频噪声。这使得模型能够专注于脑活动而非外部干扰。然后,对信号应用 Z 分数归一化,以帮助模型更好地泛化到不同患者,并降低不同患者间信号幅度的差异。
然后对脑电图记录进行加窗处理,将连续记录分割成离散状态,以便输入到我们的时空模型中。每个窗口由两个参数决定:窗口大小和步长。窗口大小决定每个窗口的长度,而步长决定每个窗口起始点之间的间隔。研究发现,3 到 5 秒之间、重叠率约为 50% 的较长窗口更为有效,这可能是因为它们具有更丰富的特征(更高的变异性),并且由于重叠,模型能够学习窗口之间的关系。

对脑电信号进行预处理滑动窗口,用于提取节点特征。
从原始信号中提取的第一个特征是每个窗口和通道的功率谱密度(PSD)。不同的脑电波频率与不同的脑活动相关,功率谱密度已被证明是癫痫发作活动的可靠预测指标(Liu S. et al.)。我们使用Welch方法计算每个窗口的归一化和对数变换后的PSD,并将delta、theta、alpha、beta和gamma频段作为特征添加进去。

用于特征的示例功率谱密度带
接下来添加的衍生特征是Hjorth移动性和复杂性参数。Lui等人发现这些特征能够预测癫痫发作是长发作还是短发作。这些特征常用于脑电图(EEG)处理,并使用mne-feature的extract_features函数以及其他统计特征进行计算。移动性参数捕捉信号的振荡,而复杂性参数衡量信号与纯正弦波的相似度。癫痫发作期间,大脑活动会出现具有不规则模式的尖峰,这些尖峰可以通过Hjorth参数以及包括均值、标准差、偏度、峰度和线段长度在内的统计指标来捕捉。将所有这些特征添加到每个通道(节点)的窗口后,我们就可以开始构建图了。


计算节点特征与癫痫发作持续时间的关联矩阵

与癫痫发作持续时间最相关的10个特征
图的构建
脑电图电极放置在头皮上时会形成自然的网格图案,但这并不是脑电信号唯一可能的图形结构。
构建的第一个图利用了头皮上电极的物理和解剖学邻近关系,以及它们所测量的脑叶。CHB-MIT 数据集使用代表两个电极的双极通道(例如,FP1-F7 和 F7-T7)。最简单的图可以通过将共享一个电极的两个通道连接起来构建。这被称为解剖邻接矩阵。
第二种图构建方法利用信号间的相关性,如果相关值高于某个阈值,则连接节点。在我们的模型中,我们使用每对通道的相位锁定值(PLV)来衡量两个通道之间的同步性,如果PLV值高于0.35(通过实验确定),则连接这两个通道。PLV是一种有效的通道连接指标,它能更精确地模拟癫痫在大脑中的传播。癫痫发作通常始于特定区域,并向整个大脑传播。因此,如果两个通道高度相关,则癫痫活动很可能继续在这些区域传播,并且在癫痫发作末期,一个通道的减缓可能意味着另一个通道的减缓。我们通过创建一个融合这两个矩阵的第三个邻接矩阵,并分别测试每个矩阵,使模型能够学习哪种表示方法更适用于癫痫发作持续时间。
相位锁定值邻接矩阵的代码:
for i, sample inenumerate(all_samples):
data = sample['normalized_segment'] # (n_ch, n_samples)
n_ch = data.shape[0]
# Compute analytic signal and phases
analytic = hilbert(data, axis=1)
phases = np.angle(analytic) # (n_ch, n_samples)
# Vectorized PLV computation using broadcasting
# phase_diff[i,j,:] = phases[i,:] - phases[j,:]
phase_diff = phases[:, np.newaxis, :] - phases[np.newaxis, :, :] # (n_ch, n_ch, n_samples)
# Phase locking value
plv = np.abs(np.mean(np.exp(1j * phase_diff), axis=-1)).astype(np.float32) # (n_ch, n_ch)
plv = np.nan_to_num(plv, nan=0.0)
plv_adj = (plv > plv_threshold).astype(np.float32)
np.fill_diagonal(plv_adj, 1.0)
sample['plv'] = plv
sample['plv_adjacency'] = plv_adj
患者 1 的邻接矩阵比较
我们模型的一个关键创新之处在于引入了虚拟节点。虚拟节点是一种图增强技术,它通过预先设计与其他节点的连接,向图中添加新节点。它们可以帮助更快地在节点(通道)之间传播数据。我们假设,引入虚拟节点可以帮助捕捉癫痫发作期间更深层的三维脑活动,这些活动无法通过头皮脑电图(EEG)捕捉到,其作用类似于植入大脑深处的立体脑电图(SEEG)电极。我们使用了两种虚拟节点设计,将初始节点特征设置为与其相连节点的平均值。
第一种方法是使用连接到所有通道的全局虚拟节点。即使癫痫发作区域在大脑中距离更近,即使头皮上相邻的区域距离较远,这种方法也能帮助快速地将癫痫发作信息传播到所有通道。
第二种方法是基于脑叶的分层设计。在这种设计中,如果双极通道包含映射到相应脑叶(额叶、颞叶、枕叶、顶叶和中央叶)的电极,则该通道首先连接到相应的脑叶虚拟节点。最后,所有脑叶虚拟节点连接到一个全局节点。这种设计更接近于大脑的实际结构。

标准10-20个脑电电极放置的虚拟节点配置
数据集拆分
模型训练采用 70/15/15 的训练集、验证集和测试集划分比例。划分数据时,必须防止同一患者同时出现在训练集和验证集/测试集中,避免数据泄露。同一患者出现在多个数据集中会降低模型对新患者的泛化能力。在急诊室或癫痫监测病房,医护人员不太可能之前见过该患者,或者观察到足够多的癫痫发作数据来预测其发作情况。因此,划分数据集时,我们确保训练集、验证集和测试集中不包含任何重复的患者数据。以往关于癫痫发作预测的研究发现,针对特定患者的模型更为准确(Sanjay Balaji 等人)。我们认为这种必要且有效的划分方式是出于任务难度的考虑。
模型架构
我们考虑的模型架构包括:一个不使用图结构、时间顺序或虚拟节点的基线多层感知器(MLP)模型;一个时空图卷积网络(GCN)模型,其空间编码采用GCN,时间编码采用1D卷积神经网络(CNN),最后连接一个MLP;一个时空图自适应树突触(GAT)模型,其架构与基于GCN的模型相同,但空间编码采用GATv2;一个GCN+长短期记忆网络(LSTM)模型,其与之前的GCN+1D CNN模型相同,但空间编码采用LSTM;以及一个GATv2+LSTM模型,其与之前的GATv2+1D CNN模型相同,但空间编码采用LSTM。这些架构的灵感来源于SGSTAN模型,该模型在癫痫发作起始时间预测方面表现良好,但尚未应用于持续时间预测。
基准 MLP 模型:
基准 MLP 模型用作与其他模型进行比较的基准,它只是使用多个 MLP 层对跨通道的节点特征进行平均(没有图结构)。
时空GCN模型:
在这个模型中,我们结合了用于时空编码的图卷积网络(GCN)、用于时间编码的一维卷积神经网络(1D CNN)以及多层感知器(MLP)回归头。GCN 和 1D CNN 都包含多个层。模型各部分的隐藏层维度和层数是需要调整的超参数。
时空 GATv2 模型:
该模型与时空 GCN 模型相同,只是我们将时空编码部分替换为 GATv2。GAT 的隐藏维度、层数和头部数是我们调整的其他超参数。
GCN+LSTM模型:
该模型与时空图卷积神经网络模型相同,只是用长短期记忆网络(LSTM)代替了1D卷积神经网络(CNN)。LSTM的隐藏层大小和层数是需要调整的额外超参数。
GATv2+LSTM模型:
该模型与时空 GATv2 模型相同,只是用 LSTM 代替了 1D CNN。LSTM 的隐藏层大小和层数是需要调整的额外超参数。

从窗口输入到持续时间预测的时空模型架构流程
最佳模型 GCN_CNN 的代码展示了空间 -> 时间 -> 回归过程:
classGCN_CNN_Model(nn.Module):
"""Spatio-Temporal GCN + CNN with Raw Signal Integration
Architecture:
1. Raw Signal Encoder: 1D CNN to embed raw EEG signals
2. Spatial Encoder (GCN): Graph convolutions over EEG channels
3. Temporal Encoder (CNN): 1D convolutions over time windows
4. Regression Head: MLP for duration prediction
"""
def__init__(self, n_features, n_samples, raw_embed_dim, gcn_hidden, gcn_layers,
cnn_channels, cnn_layers, head_hidden, dropout, virtual_mode, target_mean):
super().__init__()
self.virtual_mode = virtual_mode
self.gcn_hidden = gcn_hidden
self.raw_embed_dim = raw_embed_dim
# Raw signal encoder
self.raw_encoder = RawSignalEncoder(n_samples, raw_embed_dim, dropout)
# GNN input: features + raw embedding
gnn_input_dim = n_features + raw_embed_dim
# Spatial Encoder (GCN)
self.spatial_encoder = GCN(
in_channels=gnn_input_dim,
hidden_channels=gcn_hidden,
num_layers=gcn_layers,
out_channels=gcn_hidden,
dropout=dropout,
norm='layer_norm',
act='relu'
)
# Skip connection from raw
self.raw_skip_projection = nn.Sequential(
nn.Linear(raw_embed_dim, gcn_hidden),
nn.LayerNorm(gcn_hidden),
nn.ReLU(),
nn.Dropout(dropout)
)
# Temporal Encoder (CNN)
temporal_input_dim = gcn_hidden * 2# spatial + raw skip
self.temporal_encoder = nn.ModuleList()
in_ch = temporal_input_dim
for _ inrange(cnn_layers):
self.temporal_encoder.append(nn.Sequential(
nn.Conv1d(in_ch, cnn_channels, kernel_size=3, padding=1),
nn.BatchNorm1d(cnn_channels),
nn.ReLU(),
nn.Dropout(dropout)
))
in_ch = cnn_channels
# Regression Head
self.regression_head = MLP(channel_list=[cnn_channels, head_hidden, 1], dropout=dropout, norm='layer_norm', act='relu')
# Initialize bias to target mean for faster convergence
with torch.no_grad():
self.regression_head.lins[-1].bias.fill_(target_mean)
n_params = sum(p.numel() for p inself.parameters())
print(f'GCN+CNN params: {n_params}')
defforward(self, batch, device):
x = batch['node_features'].to(device) # (B, T, C, F)
raw = batch['raw_windows'].to(device) # (B, T, C, S)
adjacencies = batch['adjacencies']
channel_lists = batch['channels']
lengths = batch['lengths']
B, T, C, F = x.shape
S = raw.shape[-1]
raw_flat = raw.view(B * T * C, S)
raw_embed_flat = self.raw_encoder(raw_flat)
raw_embed = raw_embed_flat.view(B, T, C, self.raw_embed_dim)
raw_skip = raw_embed.mean(dim=2)
raw_skip_projected = self.raw_skip_projection(raw_skip)
combined_features = torch.cat([x, raw_embed], dim=-1)
graph_list = []
for b inrange(B):
seq_len = lengths[b].item()
adj = adjacencies[b]
channels = channel_lists[b]
# Get cached edge index (avoids redundant computation)
edge_index, _ = get_cached_edge_index(adj, channels, self.virtual_mode, device)
for t inrange(seq_len):
node_features = combined_features[b, t]
aug_features, _, _ = add_virtual_nodes(node_features, adj, channels, self.virtual_mode)
graph_list.append(Data(x=aug_features, edge_index=edge_index))
# Batch all graphs and process through GNN
batched_graphs = Batch.from_data_list(graph_list)
all_node_embeddings = self.spatial_encoder(batched_graphs.x.to(device), batched_graphs.edge_index.to(device))
window_embeddings = global_mean_pool(all_node_embeddings, batched_graphs.batch.to(device))
spatial_sequences = torch.zeros(B, T, self.gcn_hidden, device=device)
idx = 0
for b inrange(B):
seq_len = lengths[b].item()
spatial_sequences[b, :seq_len] = window_embeddings[idx:idx + seq_len]
idx += seq_len
fused_sequences = torch.cat([spatial_sequences, raw_skip_projected], dim=-1)
# Temporal encoding
temporal_input = fused_sequences.transpose(1, 2)
for conv_block inself.temporal_encoder:
temporal_input = conv_block(temporal_input)
temporal_embedding = temporal_input.mean(dim=-1)
# Regression
predictions = self.regression_head(temporal_embedding).squeeze(-1)
return predictionsgcn_cnn_cfgs = {
# Spatial encoder (GCN)
'gcn_hidden': [8, 16],
'gcn_layers': [2],
# Temporal encoder (CNN)
'cnn_ch': [8, 16],
'cnn_layers': [2],
# Optimization
'lr': [1e-4],
'dropout': [0.3],
# Graph structure
'virtual_node_mode': ['none', 'global', 'per_lobe'],
'adjacency_mode': ['anatomical', 'plv', 'combined']
}模型架构的理论依据:
在我们所有的非基线模型中,我们将空间编码器部分(使用基于图神经网络的编码器,例如 GCN 或 GATv2)与时间编码器部分(使用序列模型,例如 1D CNN 或 LSTM)相结合。这使得模型能够通过空间编码器部分恰当地捕捉电极放置的空间信息,并通过序列模型捕捉脑电波的时间信息,从而进行预测。我们将包含较简单组件(GCN 和 1D CNN)的架构与包含较复杂组件(GATv2 和 LSTM)的架构进行比较,以研究增加模型任一部分的复杂度如何影响性能。这一点至关重要,因为模型任何一部分过于复杂都可能导致过拟合,而我们希望避免这种情况。因此,我们没有尝试更复杂的组件,例如图变换器,因为我们希望先尝试更简单的模型,然后再尝试更复杂的模型,以防止过拟合。
模型训练和评估方法:
我们对每种时空架构都进行了多种超参数配置的训练,以确定最佳配置。在测试过程中,我们发现参数量远大于样本量的模型往往过拟合,在测试数据集和新患者数据上的表现均不佳。因此,我们在训练流程中评估了参数量较小的模型,以找到最佳配置。此外,我们还对邻接矩阵和虚拟节点结构的每种组合进行训练,以确定它们是否对准确率产生影响。每个模型的目标函数是每次癫痫发作持续时间的对数,这样既可以减轻异常值的影响,又无需从数据中过滤掉异常值,因为长时间的癫痫发作是需要考虑和报告的重要信息。
我们报告了三个标准指标:R²、MAE和RMSE。如果在验证集上模型性能没有提升,则在训练 10 个 epoch 后模型会提前停止。我们还尝试了不同的学习率和 dropout 率,以测试它们对模型收敛性和泛化能力(针对新患者)的影响。
任务:
我们的任务是利用各种基于图神经网络的方法预测癫痫发作的持续时间。
结果


包含提前停止策略的每个模型的最佳超参数验证曲线

性能最佳模型(GCN+CNN)的实际值和预测值。

消融研究矩阵用于虚拟节点和邻接矩阵效应
洞察
由于我们的数据集中只有175个癫痫发作样本,因此很难在不发生过拟合的情况下捕捉到癫痫发作的复杂模式。模型在验证集上的表现优于预测均值,但在测试集上的表现均为负值。这凸显了癫痫脑电图读数的复杂性以及不同患者间动态的差异。测试集和验证集中的患者数量较少,可能具有非常不同的癫痫发作模式,而这些模式由于训练集缺乏多样性而无法从中学习。同样,癫痫发作持续时间模式可能不具有普遍性,但这需要使用更大的数据集进行进一步验证。在早期使用Transformer对不同模型进行图和时间窗口的实验中,这些模型由于几乎完全记住了训练数据而严重过拟合,对测试数据集几乎没有预测能力。在比较不同的时空模型时,这种趋势仍然存在,对于小数据集而言,更简单的模型更受欢迎。
时空模型的性能优于单层或双层多层感知器(MLP),但这不太可能是由于图结构本身,而更可能是由于更复杂的时序模型所致。不同的邻接方法并未显示出显著的性能提升,由此得出结论:图编码的价值有限,甚至可能没有价值。考虑到大脑的结构以及我们对不同局灶性癫痫的理解,这一结果令人惊讶。基于相位锁定值的模型性能略优于基于电极位置的模型,这并不意外,因为临床医生放置电极的方式存在差异,而且不同患者的解剖结构也存在差异。
虚拟节点并未对任何模型带来持续的益处。无论是单个全局虚拟节点还是基于脑叶的虚拟节点,都未能提升模型性能,这表明虚拟节点不太可能代表更深层的脑电活动,也难以帮助模型理解癫痫发作的传播。在未来设计虚拟节点时,基于信号相位锁定值创建节点可能比基于脑叶的虚拟节点更有益,因为它在邻接矩阵方面略有改进,并且可能通过捕捉每个信号对其他信号的影响,更好地表示癫痫发作的三维传播。
对未来工作的建议
我们的分析存在一些已知的局限性。尽管由于收集了大量的脑电图记录,CHB-MIT 数据集在许多研究中被认为是全面的,但我们的癫痫发作水平任务最终只得到了 198 个有效样本。此外,该数据集的样本分布相当不平衡,只有一个“较长”的癫痫发作样本。另一方面,大多数癫痫发作的持续时间都集中在一个非常小的范围内,这增加了回归任务的难度。最后,CHB-MIT 数据采用双极导联方式记录,这种方式通常会掩盖全局活动,而更倾向于局部细节。这可能会对我们的图表示和最终结果产生影响。
我们的研究结果也为未来的研究指明了方向。针对数据集的限制,目前存在一些规模更大的癫痫发作脑电图数据集,但尚未公开,因此,使用这些数据集重复我们的方法可能会获得更显著的结果。此外,由于样本量较小,我们发现复杂的模型经常出现过拟合现象,但更大的数据集有望缓解这一问题,因此未来的分析可以引入Transformer模型和更深层的神经网络等工具。既往文献也成功地将该方法应用于癫痫发作持续时间的分类,因此,更具代表性的数据集或许也能在该任务上取得显著成果。
最后,具有重要临床意义的一个领域是患者特异性模型。我们观察到每位患者的癫痫发作持续时间分布存在显著差异,因此,开发能够反映患者差异的模型可以极大地改善患者的治疗效果。
结论
在本项目中,我们探索了时空图神经网络(GNN)在癫痫发作持续时间预测任务中的应用,利用脑电图(EEG)数据和图拓扑结构来模拟大脑的连接性(功能性和解剖性)。我们将空间编码器与时间序列整合到建模流程中,并尝试使用虚拟节点增强模型,以期捕捉癫痫发作建模的长程复杂性。然而,这些改进表明,模型难以泛化到未见过的数据,更简单的模型往往优于基于复杂时空图构建的模型。具体而言,引入虚拟节点带来的性能提升微乎其微,这表明虚拟节点无法模拟更深层的大脑结构。此外,患者间差异较大也对泛化模型构成重大挑战,尤其是在训练数据集较小的情况下,验证集和测试集之间的性能差距尤为突出。鉴于这些结果,我们建议未来在该领域的研究需要更大的数据集,以缓解过拟合问题,并更好地利用这些日益复杂的深度学习工具的强大功能。
参考文献
Kerr, W. T., McFarlane, K. N., & Figueiredo Pucci, G. (2024). The present and future of seizure detection, prediction, and forecasting with machine learning, including the future impact on clinical trials. Frontiers in neurology, 15, 1425490. https://doi.org/10.3389/fneur.2024.1425490
Liu, S., Wang, J., Li, S., & Cai, L. (2023). Epileptic Seizure Detection and Prediction in EEGs Using Power Spectra Density Parameterization. IEEE transactions on neural systems and rehabilitation engineering : a publication of the IEEE Engineering in Medicine and Biology Society, 31, 3884–3894. https://doi.org/10.1109/TNSRE.2023.3317093
Liu, Y., Razavi Hesabi, Z., Cook, M., & Kuhlmann, L. (2022). Epileptic seizure onset predicts its duration. European journal of neurology, 29(2), 375–381. https://doi.org/10.1111/ene.15166
Sanjay Balaji, S., Zhang, Z., Sha, Z., Henry, T. R., & Parhi, K. K. (2025). Patient-specific long-term seizure prediction via multi-model classification. Journal of neural engineering, 22(6), 10.1088/1741–2552/ae1875. https://doi.org/10.1088/1741-2552/ae1875



内容中包含的图片若涉及版权问题,请及时与我们联系删除





评论
沙发等你来抢