因果感知KAN多模态驾驶疲劳融合:EEG+EOG+ECG因果建模+Kolmogorov-Arnold网络

论文信息

项目 内容
标题 Multi-modal driving fatigue detection with causal-aware fusion and Kolmogorov-Arnold network
期刊 Information Fusion, Vol. 133, Article 104255
发表 2026年2月26日
作者 Xu Xu, Ghulam Muhammad
链接 https://www.sciencedirect.com/science/article/abs/pii/S156625352600134X
核心方法 因果感知融合 + KAN层
输入 EEG + EOG + ECG 多模态生理信号

核心创新

  1. 因果感知融合:不简单拼接特征,而是建模模态间因果关系
  2. KAN融合层:用Kolmogorov-Arnold Network替代传统MLP进行跨模态融合
  3. 交通安全性驱动:以交通安全为目标的疲劳检测框架
  4. 因果图建模:EEG→认知疲劳→EOG→眨眼行为→ECG→自主神经反应

问题定义

传统多模态融合的局限

融合方法 做法 局限
早期融合 原始信号拼接 噪声叠加
晚期融合 各模态独立决策→投票 丢失交互信息
注意力融合 学习模态权重 忽略因果关系
因果融合 建模模态间因果图 本论文

生理信号间的因果关系

flowchart TD
    A[EEG: 认知疲劳] -->|因果| B[EOG: 眨眼/扫视]
    A -->|因果| C[ECG: 心率变异]
    B -->|反映| D[疲劳状态]
    C -->|反映| D
    A -->|直接反映| D
    
    E[驾驶时长] -->|导致| A
    F[环境因素] -->|影响| A
    F -->|影响| B
因果路径 描述
EEG→EOG 认知疲劳导致眨眼频率/时长变化
EEG→ECG 中枢疲劳影响自主神经系统→心率变异
ECG→EOG 自主神经反应影响眼动行为
驾驶时长→EEG 时间累积导致认知疲劳

方法详解

因果感知融合+KAN架构

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
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
import torch
import torch.nn as nn
import numpy as np

class CausalDiscoveryModule(nn.Module):
"""
因果发现模块

从多模态数据中学习模态间因果关系
使用注意力掩码学习因果图
"""

def __init__(self, n_modalities: int = 3,
embed_dim: int = 64):
super().__init__()
self.n_modalities = n_modalities

# 因果掩码(可学习邻接矩阵)
self.causal_mask = nn.Parameter(
torch.randn(n_modalities, n_modalities) * 0.1
)

# 模态嵌入
self.modality_embed = nn.ModuleList([
nn.Linear(1, embed_dim) for _ in range(n_modalities)
])

def forward(self, modalities: list) -> tuple:
"""
Args:
modalities: [eeg, eog, ecg] 每个[B, T, 1]

Returns:
fused: [B, embed_dim] 融合特征
causal_graph: [M, M] 因果图
"""
# 模态嵌入
embeddings = []
for i, m in enumerate(modalities):
# 每个模态用对应的嵌入层
m_embed = self.modality_embed[i](m) # [B, T, embed_dim]
# 时序池化
m_embed = m_embed.mean(dim=1) # [B, embed_dim]
embeddings.append(m_embed)

# 堆叠 [B, M, embed_dim]
stacked = torch.stack(embeddings, dim=1)

# 因果图(softmax归一化)
causal_graph = torch.softmax(self.causal_mask, dim=-1)

# 因果传播:G_sigmoid(A) × X
# 每个模态的表示被其因果父节点更新
updated = torch.einsum(
'mm,bme->bme',
causal_graph,
stacked
)

# 融合
fused = updated.mean(dim=1) # [B, embed_dim]

return fused, causal_graph


class KANFusionLayer(nn.Module):
"""
KAN融合层

使用Kolmogorov-Arnold Network进行跨模态融合
替代传统MLP,更少参数更高表达力
"""

def __init__(self, in_dim: int, out_dim: int,
grid_size: int = 5):
super().__init__()

# KAN边:可学习样条函数
self.base_linear = nn.Linear(in_dim, out_dim)

# 样条参数(B样条基函数权重)
self.spline_weight = nn.Parameter(
torch.randn(out_dim, in_dim, grid_size + 3) * 0.1
)

# 网格
h = 1.0 / grid_size
grid = torch.linspace(-3*h, 1+3*h, grid_size + 7)
self.register_buffer('grid', grid)

def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
KAN前向传播

output = base_linear(x) + spline(x)
"""
base = self.base_linear(x)

# 简化:用sigmoid近似B样条
x_expanded = x.unsqueeze(-1) # [B, in_dim, 1]
grid = self.grid.view(1, 1, -1) # [1, 1, grid_size+7]

# B样条基(简化版)
b = torch.sigmoid((x_expanded - grid) * grid_size)
b = b[..., :self.spline_weight.shape[-1]]

# 样条部分
spline = torch.einsum(
'big,oig->bog',
b.expand(x.shape[0], -1, -1),
self.spline_weight
)

return base + spline


class CausalKANFatigueModel(nn.Module):
"""
因果感知KAN多模态疲劳检测模型

架构:
1. 各模态特征提取(EEG/EOG/ECG)
2. 因果发现模块(学习模态间因果关系)
3. KAN融合层(因果传播后融合)
4. KAN分类器(疲劳/清醒)
"""

def __init__(self, n_classes: int = 3):
super().__init__()

# 各模态特征提取
self.eeg_encoder = nn.Sequential(
nn.Conv1d(1, 32, 7, stride=2),
nn.BatchNorm1d(32),
nn.ReLU(),
nn.Conv1d(32, 64, 5, stride=2),
nn.BatchNorm1d(64),
nn.ReLU(),
nn.AdaptiveAvgPool1d(1),
nn.Flatten(),
)
self.eog_encoder = nn.Sequential(
nn.Conv1d(1, 32, 5, stride=2),
nn.BatchNorm1d(32),
nn.ReLU(),
nn.Conv1d(32, 64, 3, stride=2),
nn.BatchNorm1d(64),
nn.ReLU(),
nn.AdaptiveAvgPool1d(1),
nn.Flatten(),
)
self.ecg_encoder = nn.Sequential(
nn.Conv1d(1, 32, 15, stride=4),
nn.BatchNorm1d(32),
nn.ReLU(),
nn.Conv1d(32, 64, 7, stride=2),
nn.BatchNorm1d(64),
nn.ReLU(),
nn.AdaptiveAvgPool1d(1),
nn.Flatten(),
)

# 因果发现
self.causal = CausalDiscoveryModule(n_modalities=3)

# KAN融合
self.kan_fusion = KANFusionLayer(64, 128, grid_size=5)

# KAN分类器
self.classifier = KANFusionLayer(128, n_classes, grid_size=5)

def forward(self, eeg, eog, ecg):
"""
Args:
eeg: [B, 1, T_eeg]
eog: [B, 1, T_eog]
ecg: [B, 1, T_ecg]

Returns:
logits: [B, n_classes]
causal_graph: [3, 3]
"""
# 1. 特征提取
eeg_feat = self.eeg_encoder(eeg).unsqueeze(1) # [B, 1, 64]
eog_feat = self.eog_encoder(eog).unsqueeze(1)
ecg_feat = self.ecg_encoder(ecg).unsqueeze(1)

# 2. 因果发现
fused, causal_graph = self.causal(
[eeg_feat, eog_feat, ecg_feat]
)

# 3. KAN融合
fused = self.kan_fusion(fused)
fused = torch.relu(fused)

# 4. 分类
logits = self.classifier(fused)

return logits, causal_graph


# 测试
if __name__ == "__main__":
model = CausalKANFatigueModel(n_classes=3)

# 模拟输入
batch = 4
eeg = torch.randn(batch, 1, 200) # 200采样点
eog = torch.randn(batch, 1, 200)
ecg = torch.randn(batch, 1, 500) # ECG采样率更高

logits, causal = model(eeg, eog, ecg)

print(f"输出: {logits.shape}")
print(f"因果图:\n{causal.detach()}")

# 模态标签
modality_names = ['EEG', 'EOG', 'ECG']
print(f"\n因果路径强度:")
for i in range(3):
for j in range(3):
if i != j:
print(f" {modality_names[i]}{modality_names[j]}: "
f"{causal[i,j].item():.3f}")

total_params = sum(p.numel() for p in model.parameters())
print(f"\n总参数: {total_params:,}")

实验结果

性能对比

方法 融合方式 准确率 F1 参数量
EEG only - 78.5% 77.8% 120K
EOG only - 72.3% 71.5% 120K
ECG only - 68.7% 67.9% 120K
拼接融合 早期 85.2% 84.6% 380K
注意力融合 学习权重 88.7% 88.1% 420K
因果+KAN 因果+KAN 93.5% 93.1% 285K

因果图学习结果

因果路径 强度 生理意义
EEG→EOG 0.82 认知疲劳→眨眼变化
EEG→ECG 0.71 中枢疲劳→心率变异
ECG→EOG 0.45 自主神经→眼动
EOG→EEG 0.12 弱反向因果
EOG→ECG 0.08 弱反向因果
ECG→EEG 0.15 弱反向因果

KAN vs MLP融合对比

融合层 准确率 参数量 可解释性
MLP融合 88.7% 420K
KAN融合 93.5% 285K
提升 +4.8% -32%

IMS开发启示

1. 因果感知融合对IMS的价值

价值 描述
可解释性 告诉开发者”EEG→EOG是主要因果路径”
鲁棒性 一个模态故障时因果图自动降权
效率 KAN比MLP少32%参数,更适合边缘
因果发现 无需先验因果知识,从数据中学习

2. 完整多模态管道

模态 传感器 编码器 因果角色
EEG 耳道电极 Conv1D 源头因果
EOG DMS摄像头眼动 Conv1D 中间效应
ECG rPPG/胸带 Conv1D 中间效应
融合 - 因果+KAN -

3. 与已有管道的集成

组件 来源 与因果KAN的关系
DeltaGateNet EEG编码 (#24) EEG特征提取 可作为eeg_encoder
rPPG心率信号 (#14, #03) 心率→ECG模态 替代ECG传感器
眼动PERCLOS DMS摄像头 EOG模态来源
因果KAN融合 (本论文) 多模态融合 顶层融合引擎

总结

因果感知KAN多模态融合是疲劳检测的高级融合方案:

  1. 因果发现:从数据中学习EEG→EOG→ECG因果图,非简单拼接
  2. KAN融合:比MLP少32%参数,准确率高4.8%(93.5% vs 88.7%)
  3. 可解释性:因果图揭示”认知疲劳→眨眼变化”为主要路径
  4. 模态鲁棒:因果图自动降权故障模态
  5. IMS顶层引擎:整合EEG/EOG/ECG三路信号的因果融合决策层

https://dapalm.com/2026/09/22/2026-09-22-06-causal-kan-multimodal-fatigue-eeg-eog-ecg-ims/
作者
Mars
发布于
2026年9月22日
许可协议