KAN-CLUE:Kolmogorov-Arnold网络+不确定性感知持续学习TinyML驾驶员疲劳检测

论文信息

项目 内容
标题 Uncertainty-Aware Continual TinyML Driver Fatigue Detection with Kolmogorov–Arnold Networks at the IoT Edge
期刊 Applied System Innovation (MDPI), Vol. 9, Issue 7, Article 147
发表 2026年7月8日
链接 https://doi.org/10.3390/asi9070147
核心方法 KAN-CLUE = CNN骨干 + Kolmogorov-Arnold Network + 不确定性量化 + 持续学习
输入 近红外眼周图像
部署 IoT边缘TinyML设备

核心创新

  1. KAN替代MLP:在边缘设备上用Kolmogorov-Arnold Network替代传统多层感知机,更少参数更高精度
  2. 不确定性感知:模型输出预测置信度,低置信度时不触发警报,降低误报
  3. 持续学习:新驾驶员数据在线增量学习,无需重新训练
  4. TinyML部署:专为资源受限IoT边缘设备设计(<100KB RAM)

问题定义

边缘疲劳检测三难选择

维度 传统方案 KAN-CLUE方案
精度 CNN+MLP高但参数多 KAN精度相当参数更少
适应性 固定模型→新驾驶员性能下降 持续学习→适应新用户
可靠性 过度自信→误报 不确定性量化→降级输出

Kolmogorov-Arnold Network (KAN) 原理

flowchart LR
    subgraph MLP[传统MLP]
        A1[输入] --> B1[Linear: Wx+b]
        B1 --> C1[激活σ]
        C1 --> D1[Linear: Wx+b]
        D1 --> E1[输出]
    end
    
    subgraph KAN[KAN]
        A2[输入] --> B2[可学习样条函数φ]
        B2 --> C2[求和]
        C2 --> D2[可学习样条函数φ]
        D2 --> E2[输出]
    end
特性 MLP KAN
激活位置 节点(固定) 边(可学习)
函数类型 固定ReLU/Sigmoid B样条曲线(可学习)
参数效率 中等 高(少参数高性能)
可解释性
边缘适用 一般 优秀

方法详解

KAN-CLUE架构

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
import torch
import torch.nn as nn
import numpy as np

class KANLayer(nn.Module):
"""
Kolmogorov-Arnold Network层

核心区别:可学习激活函数在边上而非节点上
使用B样条逼近任意一维函数
"""

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

# 基函数权重(线性部分)
self.base_weight = nn.Parameter(
torch.randn(out_dim, in_dim) * 0.1
)

# 样条权重(非线性部分)
self.spline_weight = nn.Parameter(
torch.randn(out_dim, in_dim, grid_size + spline_order) * 0.1
)

# 样条网格点
h = 1.0 / grid_size
grid = torch.arange(-spline_order, grid_size + spline_order + 1) * h
self.register_buffer('grid', grid)

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

Args:
x: [batch, in_dim]

Returns:
output: [batch, out_dim]
"""
batch_size = x.shape[0]

# 基函数(线性)
base = torch.einsum('bi,oi->bo', x, self.base_weight)

# B样条基函数
x_expanded = x.unsqueeze(-1) # [batch, in_dim, 1]
grid_expanded = self.grid.view(1, 1, -1) # [1, 1, grid_size+2*spline_order+1]

# Cox-de Boor递推计算B样条基
b_splines = self._bspline_basis(
x_expanded.expand(-1, -1, grid_expanded.shape[-1]),
grid_expanded.expand(x.shape[0], -1, -1)
)

# 样条部分
spline = torch.einsum(
'big,oig->bog',
b_splines,
self.spline_weight
)

return base + spline

def _bspline_basis(self, x: torch.Tensor,
grid: torch.Tensor) -> torch.Tensor:
"""计算B样条基函数值(简化版)"""
order = self.spline_order
if order == 0:
return ((x >= grid[..., :-1]) &
(x < grid[..., 1:])).float()

# 递推
B_prev = self._bspline_basis(x, grid)
if B_prev.shape[-1] > 1:
denom1 = grid[..., order:B_prev.shape[-1]] - grid[..., :-order]
denom1 = torch.where(denom1 != 0, denom1,
torch.ones_like(denom1))
term1 = (x - grid[..., :-order]) / denom1 * B_prev[..., :-1]

denom2 = grid[..., order+1:] - grid[..., 1:B_prev.shape[-1]]
denom2 = torch.where(denom2 != 0, denom2,
torch.ones_like(denom2))
term2 = (grid[..., order+1:] - x) / denom2 * B_prev[..., 1:]

return term1 + term2

return B_prev


class KANCLUEModel(nn.Module):
"""
KAN-CLUE: CNN骨干 + KAN分类器 + 不确定性量化

架构:
1. 轻量CNN提取眼周图像特征
2. KAN层替代MLP进行分类
3. 蒙特卡洛Dropout量化不确定性
"""

def __init__(self, num_classes: int = 3,
feature_dim: int = 128,
kan_grid: int = 5):
super().__init__()

# CNN骨干(眼周图像特征提取)
self.backbone = nn.Sequential(
nn.Conv2d(1, 16, 3, stride=2, padding=1),
nn.BatchNorm2d(16),
nn.ReLU6(),
nn.Conv2d(16, 32, 3, stride=2, padding=1),
nn.BatchNorm2d(32),
nn.ReLU6(),
nn.Conv2d(32, 64, 3, stride=2, padding=1),
nn.BatchNorm2d(64),
nn.ReLU6(),
nn.AdaptiveAvgPool2d(1),
)
self.flatten = nn.Flatten()

# KAN分类器(替代MLP)
self.kan1 = KANLayer(64, feature_dim, grid_size=kan_grid)
self.kan2 = KANLayer(feature_dim, num_classes, grid_size=kan_grid)

# 不确定性量化(MC Dropout)
self.dropout = nn.Dropout(0.1)

def forward(self, x: torch.Tensor,
n_samples: int = 1) -> tuple:
"""
前向传播

Args:
x: [batch, 1, H, W] 近红外眼周图像
n_samples: MC采样次数(>1时输出不确定性)

Returns:
logits: [batch, num_classes]
uncertainty: [batch] 预测熵
"""
features = self.backbone(x)
features = self.flatten(features)

# MC Dropout采样
logits_list = []
for _ in range(n_samples):
h = self.dropout(features)
h = self.kan1(h)
h = torch.relu(h)
logits = self.kan2(h)
logits_list.append(logits)

logits_stack = torch.stack(logits_list) # [n_samples, batch, classes]

if n_samples > 1:
# 平均预测
avg_logits = logits_stack.mean(dim=0)
# 预测熵(不确定性)
probs = torch.softmax(avg_logits, dim=-1)
entropy = -torch.sum(
probs * torch.log(probs + 1e-8), dim=-1
)
return avg_logits, entropy
else:
return logits_stack[0], torch.zeros(x.shape[0])


# 测试
if __name__ == "__main__":
model = KANCLUEModel(num_classes=3) # 清醒/疲劳/嗜睡

# 模拟眼周图像 [batch, 1, 48, 48]
x = torch.randn(4, 1, 48, 48)

# 单次推理
logits, _ = model(x, n_samples=1)
print(f"单次推理 logits: {logits.shape}")

# 不确定性推理(10次MC采样)
logits, uncertainty = model(x, n_samples=10)
print(f"不确定性推理 logits: {logits.shape}")
print(f"不确定性: {uncertainty}")
print(f"高置信度(低熵): {(uncertainty < 0.5).sum()}/4")

# 参数量对比
total_params = sum(p.numel() for p in model.parameters())
print(f"\nKAN-CLUE总参数: {total_params:,}")

# 与MLP对比
mlp_model = nn.Sequential(
nn.Conv2d(1, 16, 3, stride=2, padding=1),
nn.ReLU(),
nn.Conv2d(16, 32, 3, stride=2, padding=1),
nn.ReLU(),
nn.Conv2d(32, 64, 3, stride=2, padding=1),
nn.ReLU(),
nn.AdaptiveAvgPool2d(1),
nn.Flatten(),
nn.Linear(64, 128),
nn.ReLU(),
nn.Linear(128, 3)
)
mlp_params = sum(p.numel() for p in mlp_model.parameters())
print(f"等效MLP参数: {mlp_params:,}")
print(f"压缩比: {mlp_params/total_params:.1f}x")

持续学习机制

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
class ContinualLearning:
"""
持续学习模块

新驾驶员数据增量学习,不遗忘旧驾驶员
"""

def __init__(self, model, lr=1e-4, ewc_lambda=0.1):
self.model = model
self.lr = lr
self.ewc_lambda = ewc_lambda # EWC正则化
self.fisher_matrix = None # 重要参数记录

def update_fisher(self, dataloader):
"""计算Fisher信息矩阵(标记重要参数)"""
fisher = {}
for name, param in self.model.named_parameters():
fisher[name] = torch.zeros_like(param)

self.model.eval()
for x, y in dataloader:
self.model.zero_grad()
logits, _ = self.model(x, n_samples=1)
loss = nn.functional.cross_entropy(logits, y)
loss.backward()

for name, param in self.model.named_parameters():
if param.grad is not None:
fisher[name] += param.grad.pow(2)

for name in fisher:
fisher[name] /= len(dataloader)

self.fisher_matrix = fisher

def continual_update(self, dataloader):
"""
持续学习更新

EWC约束:重要参数变化小,防止遗忘
"""
optimizer = torch.optim.Adam(
self.model.parameters(), lr=self.lr
)

self.model.train()
for x, y in dataloader:
optimizer.zero_grad()

logits, uncertainty = self.model(x, n_samples=5)
ce_loss = nn.functional.cross_entropy(logits, y)

# 不确定性正则化(鼓励高置信度)
unc_loss = uncertainty.mean()

# EWC正则化(防止遗忘)
ewc_loss = 0
if self.fisher_matrix:
for name, param in self.model.named_parameters():
if name in self.fisher_matrix:
ewc_loss += (
self.fisher_matrix[name] *
param.pow(2)
).sum()

total_loss = ce_loss + 0.01 * unc_loss + \
self.ewc_lambda * ewc_loss

total_loss.backward()
optimizer.step()

实验结果

性能对比

方法 准确率 参数量 模型大小 推理延迟 不确定性
CNN+MLP 89.2% 320K 1.2MB 12ms
CNN+MLP+MC Dropout 88.8% 320K 1.2MB 120ms
FastKAN-DDD 91.5% 180K 0.7MB 8ms
KAN-CLUE 92.3% 95K 0.38MB 15ms

持续学习效果

场景 无持续学习 EWC持续学习 KAN-CLUE
驾驶员A(训练) 93.1% 93.1% 93.1%
驾驶员B(新) 71.2% 84.5% 89.7%
驾驶员A(遗忘测试) 85.3%↓ 91.2% 92.8%
10人平均 76.8% 87.3% 90.1%

不确定性量化效果

场景 置信度 误报率 漏报率
高置信度(>0.8) 2.1% 1.5%
中置信度(0.5-0.8) ⚠️ 8.3% 5.2%
低置信度(<0.5) - -
无不确定性 - 12.7% 3.8%

IMS开发启示

1. KAN在IMS中的应用价值

应用场景 KAN优势 MLP劣势
边缘部署 95K参数 320K参数
新用户适应 持续学习 重新训练
可靠报警 不确定性降级 过度自信误报
可解释性 样条函数可视化 黑盒

2. 与DeltaGateNet的对比

维度 DeltaGateNet (#24) KAN-CLUE 选择建议
输入 EEG时序 眼周图像 不同模态
参数 45K 95K DeltaGateNet更小
不确定性 KAN-CLUE更可靠
持续学习 KAN-CLUE更灵活
可解释性 中等 KAN-CLUE

3. 部署架构

组件 模型 输入 延迟
眼周疲劳 KAN-CLUE NIR眼周图像 15ms
EEG疲劳 DeltaGateNet 耳道EEG 2ms
融合决策 不确定性加权 两路输出 1ms
总计 - - 18ms

总结

KAN-CLUE代表了驾驶员疲劳检测的新一代边缘AI:

  1. KAN替代MLP:95K参数实现92.3%准确率,比MLP小3.4倍
  2. 不确定性量化:MC Dropout输出预测熵,低置信度时降级而非误报
  3. 持续学习:EWC机制使新驾驶员准确率从71.2%提升到89.7%
  4. TinyML就绪:0.38MB模型+15ms推理,适用于MCU级边缘设备
  5. 与DeltaGateNet互补:眼周图像+EEG双模态融合,不确定性加权决策

https://dapalm.com/2026/09/22/2026-09-22-04-kan-clue-tinyml-fatigue-uncertainty-continual-ims/
作者
Mars
发布于
2026年9月22日
许可协议