Transformer 实时疲劳检测:Nature 2025 论文深度解读与代码复现

Transformer 实时疲劳检测:Nature 2025 论文深度解读与代码复现

一、论文信息

标题: Real-time driver drowsiness detection using transformer architectures: a novel deep learning approach
期刊: Scientific Reports (Nature), 2025
DOI: 10.1038/s41598-025-02111-x
发表时间: 2025年5月


二、研究背景

2.1 传统疲劳检测方法的局限

方法 局限性
PERCLOS 仅依赖眼睑开度,忽略时序特征
CNN 方法 难以建模长时间依赖关系
手工特征 泛化能力差,场景受限

2.2 Transformer 优势

  • 全局时序建模 - 自注意力机制捕捉长期依赖
  • 可扩展性 - 预训练 + 微调范式
  • 端到端学习 - 无需手工特征工程

三、方法详解

3.1 整体架构

graph TB
    A[视频输入] --> B[帧采样]
    B --> C[特征提取 ViT]
    
    C --> D[时序编码]
    D --> E[Transformer Encoder]
    
    E --> F[疲劳分类头]
    F --> G{疲劳等级}
    
    G -->|Level 0| H[正常]
    G -->|Level 1| I[轻度疲劳]
    G -->|Level 2| J[重度疲劳]

3.2 视觉特征提取(ViT)

ViT-B/16 配置:

  • 输入图像:224×224 RGB
  • Patch 大小:16×16
  • Patch 数量:196
  • 嵌入维度:768
  • 层数:12
  • 参数量:86M
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
import torch
import torch.nn as nn
import math

class PatchEmbedding(nn.Module):
"""
图像 → Patch 嵌入
"""

def __init__(self, img_size=224, patch_size=16, in_channels=3, embed_dim=768):
super().__init__()

self.img_size = img_size
self.patch_size = patch_size
self.num_patches = (img_size // patch_size) ** 2

# 卷积实现 Patch 切分
self.proj = nn.Conv2d(
in_channels, embed_dim,
kernel_size=patch_size, stride=patch_size
)

def forward(self, x):
"""
Args:
x: 输入图像, shape=(B, C, H, W)

Returns:
embeddings: Patch 嵌入, shape=(B, num_patches, embed_dim)
"""
# (B, C, H, W) -> (B, embed_dim, H/P, W/P)
x = self.proj(x)

# Flatten: (B, embed_dim, num_patches)
x = x.flatten(2)

# Transpose: (B, num_patches, embed_dim)
x = x.transpose(1, 2)

return x


class MultiHeadAttention(nn.Module):
"""
多头自注意力机制
"""

def __init__(self, embed_dim=768, num_heads=12):
super().__init__()

self.num_heads = num_heads
self.head_dim = embed_dim // num_heads
self.scale = self.head_dim ** -0.5

# Q, K, V 线性层
self.qkv = nn.Linear(embed_dim, embed_dim * 3, bias=True)
self.proj = nn.Linear(embed_dim, embed_dim)

def forward(self, x):
"""
Args:
x: 输入序列, shape=(B, N, D)

Returns:
output: 注意力输出, shape=(B, N, D)
"""
B, N, D = x.shape

# 计算 Q, K, V
qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim)
qkv = qkv.permute(2, 0, 3, 1, 4) # (3, B, H, N, head_dim)
q, k, v = qkv[0], qkv[1], qkv[2]

# 注意力分数
attn = (q @ k.transpose(-2, -1)) * self.scale # (B, H, N, N)
attn = attn.softmax(dim=-1)

# 注意力加权
x = (attn @ v).transpose(1, 2).reshape(B, N, D)

# 输出投影
x = self.proj(x)

return x


class TransformerBlock(nn.Module):
"""
Transformer 编码器块
"""

def __init__(self, embed_dim=768, num_heads=12, mlp_ratio=4.0):
super().__init__()

# Layer Normalization
self.norm1 = nn.LayerNorm(embed_dim)
self.norm2 = nn.LayerNorm(embed_dim)

# Multi-Head Attention
self.attn = MultiHeadAttention(embed_dim, num_heads)

# MLP
self.mlp = nn.Sequential(
nn.Linear(embed_dim, int(embed_dim * mlp_ratio)),
nn.GELU(),
nn.Linear(int(embed_dim * mlp_ratio), embed_dim)
)

def forward(self, x):
# Pre-norm + Attention
x = x + self.attn(self.norm1(x))

# Pre-norm + MLP
x = x + self.mlp(self.norm2(x))

return x


class FatigueTransformer(nn.Module):
"""
Transformer 疲劳检测模型
"""

def __init__(self, num_frames=16, embed_dim=768, num_heads=12, num_layers=12, num_classes=3):
super().__init__()

self.num_frames = num_frames

# 1. Patch Embedding
self.patch_embed = PatchEmbedding(
img_size=224, patch_size=16,
in_channels=3, embed_dim=embed_dim
)

# 2. 位置编码
self.pos_embed = nn.Parameter(
torch.zeros(1, self.patch_embed.num_patches + 1, embed_dim)
)
nn.init.trunc_normal_(self.pos_embed, std=0.02)

# 3. Class Token
self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
nn.init.trunc_normal_(self.cls_token, std=0.02)

# 4. Transformer Blocks
self.blocks = nn.ModuleList([
TransformerBlock(embed_dim, num_heads)
for _ in range(num_layers)
])

# 5. Layer Norm
self.norm = nn.LayerNorm(embed_dim)

# 6. 分类头
self.head = nn.Linear(embed_dim, num_classes)

def forward(self, x):
"""
Args:
x: 视频帧序列, shape=(B, T, C, H, W)

Returns:
logits: 疲劳等级, shape=(B, num_classes)
"""
B, T, C, H, W = x.shape

# 逐帧提取特征
frame_features = []
for t in range(T):
frame = x[:, t] # (B, C, H, W)

# Patch embedding
patches = self.patch_embed(frame) # (B, num_patches, D)

# 添加 class token
cls_tokens = self.cls_token.expand(B, -1, -1)
patches = torch.cat([cls_tokens, patches], dim=1)

# 添加位置编码
patches = patches + self.pos_embed

frame_features.append(patches)

# 时序聚合(平均池化)
x = torch.stack(frame_features, dim=1) # (B, T, N+1, D)
x = x.mean(dim=1) # (B, N+1, D)

# Transformer 编码
for block in self.blocks:
x = block(x)

# Layer Norm
x = self.norm(x)

# 分类(使用 class token)
x = x[:, 0] # (B, D)
logits = self.head(x) # (B, num_classes)

return logits


# 测试代码
if __name__ == "__main__":
# 创建模型
model = FatigueTransformer(
num_frames=16,
embed_dim=768,
num_heads=12,
num_layers=12,
num_classes=3
)

# 模拟输入(16帧视频)
B = 2
video = torch.randn(B, 16, 3, 224, 224)

# 前向传播
logits = model(video)

print(f"输入形状: {video.shape}")
print(f"输出形状: {logits.shape}")
print(f"预测: {logits.argmax(dim=1)}")

# 统计参数量
total_params = sum(p.numel() for p in model.parameters())
print(f"参数量: {total_params:,} ({total_params/1e6:.1f}M)")

四、实验结果

4.1 数据集

数据集 规模 标注
Drowsy Driver Detection 10,000 视频 3类(正常/轻度/重度)
NTHU-DDD 36 受试者 PERCLOS 标注
RLDD 30 小时 疲劳等级

4.2 性能对比

方法 准确率 F1-score 推理时间
CNN-3D 91.2% 90.8% 35ms
ResNet+LSTM 93.5% 93.1% 28ms
ViT+Transformer(本文) 97.8% 97.5% 22ms

4.3 边缘部署性能

平台 FP32 延迟 INT8 延迟 功耗
QCS8255 45ms 18ms 1.5W
Jetson Orin 32ms 12ms 2.8W
RTX 4090 8ms 4ms 15W

五、关键技术分析

5.1 时序建模策略

论文采用两种策略:

  1. 帧级特征聚合(本文采用)

    • 逐帧提取 ViT 特征
    • 平均池化融合时序信息
    • 优点:计算高效,适合实时应用
  2. 时空 Transformer(替代方案)

    • 同时建模空间和时间维度
    • 计算量大(O(T²·N²))
    • 优点:更强时序建模能力

5.2 预训练策略

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
def load_pretrained_vit(model, pretrained_path='vit_b_16.pth'):
"""
加载预训练 ViT 权重
"""
# 加载 ImageNet 预训练权重
state_dict = torch.load(pretrained_path)

# 过滤不匹配的键
model_dict = model.state_dict()
pretrained_dict = {
k: v for k, v in state_dict.items()
if k in model_dict and v.shape == model_dict[k].shape
}

# 加载权重
model_dict.update(pretrained_dict)
model.load_state_dict(model_dict)

print(f"加载预训练权重: {len(pretrained_dict)}/{len(model_dict)} 层")

return model

六、IMS 集成方案

6.1 实时推理管道

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

class RealTimeFatigueDetector:
"""
实时疲劳检测管道
"""

def __init__(self, model_path='fatigue_transformer.onnx'):
import onnxruntime as ort

# 加载模型
self.session = ort.InferenceSession(model_path)

# 帧缓存
self.frame_buffer = []
self.buffer_size = 16

# 疲劳状态
self.fatigue_level = 0
self.frame_count = 0

def process_frame(self, frame):
"""
处理单帧

Args:
frame: BGR 图像, shape=(H, W, 3)

Returns:
fatigue_level: 疲劳等级 (0-2)
"""
# 1. 预处理
frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
frame_resized = cv2.resize(frame_rgb, (224, 224))
frame_normalized = (frame_resized / 255.0).astype(np.float32)

# 2. 添加到缓存
self.frame_buffer.append(frame_normalized)

# 3. 缓存未满,返回正常
if len(self.frame_buffer) < self.buffer_size:
return 0

# 4. 缓存已满,执行推理
if len(self.frame_buffer) > self.buffer_size:
self.frame_buffer.pop(0)

# 5. 组装输入
video_input = np.stack(self.frame_buffer, axis=0) # (16, H, W, 3)
video_input = video_input.transpose(3, 0, 1, 2) # (3, 16, H, W)
video_input = np.expand_dims(video_input, 0) # (1, 3, 16, H, W)

# 6. 推理
outputs = self.session.run(None, {'input': video_input})
logits = outputs[0][0] # (3,)

# 7. 解析结果
self.fatigue_level = np.argmax(logits)
confidence = np.max(logits)

# 8. 更新计数
self.frame_count += 1

return self.fatigue_level, confidence

def get_status(self):
"""
获取当前状态

Returns:
status: 状态字符串
"""
levels = ["正常", "轻度疲劳", "重度疲劳"]
return levels[self.fatigue_level]


# 使用示例
if __name__ == "__main__":
detector = RealTimeFatigueDetector()

# 模拟视频流
cap = cv2.VideoCapture(0)

while True:
ret, frame = cap.read()
if not ret:
break

# 检测
level, conf = detector.process_frame(frame)

# 显示结果
cv2.putText(
frame,
f"Fatigue: {detector.get_status()} ({conf:.2f})",
(10, 30),
cv2.FONT_HERSHEY_SIMPLEX,
1.0,
(0, 255, 0) if level == 0 else (0, 255, 255) if level == 1 else (0, 0, 255),
2
)

cv2.imshow('Fatigue Detection', frame)

if cv2.waitKey(1) & 0xFF == ord('q'):
break

cap.release()
cv2.destroyAllWindows()

6.2 性能优化

模型量化(INT8):

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
# PTQ 量化流程
import torch.quantization as quant

# 1. 准备量化配置
model.qconfig = quant.get_default_qconfig('fbgemm')

# 2. 融合层(提升量化精度)
quant.fuse_modules(model, [['blocks.0.attn.qkv']]) # 示例

# 3. 准备量化
quant.prepare(model, inplace=True)

# 4. 校准(使用真实数据)
with torch.no_grad():
for batch in calibration_loader:
model(batch)

# 5. 转换为 INT8
quant.convert(model, inplace=True)

# 6. 保存
torch.save(model.state_dict(), 'fatigue_transformer_int8.pth')

七、开发检查清单

7.1 模型训练

  • 准备数据集(≥10,000 视频)
  • 加载 ImageNet 预训练权重
  • 训练时序聚合模块
  • 验证准确率 ≥97%

7.2 边缘部署

  • 导出 ONNX 格式
  • INT8 量化校准
  • 测试推理延迟 ≤30ms
  • 验证功耗 ≤2W

7.3 场景测试

场景 测试条件 预期结果
白天 光照 500±100 lux 准确率 ≥98%
夜间 红外补光 准确率 ≥95%
逆光 强光干扰 准确率 ≥92%
遮挡 眼镜/口罩 准确率 ≥90%

八、参考资源

  1. 论文原文: https://www.nature.com/articles/s41598-025-02111-x
  2. ViT 论文: https://arxiv.org/abs/2010.11929
  3. PyTorch 实现: https://github.com/lucidrains/vit-pytorch
  4. TensorRT 部署: https://docs.nvidia.com/deeplearning/tensorrt/

九、总结

Transformer 疲劳检测实现97.8% 准确率,关键突破:

  1. 全局时序建模 - 自注意力捕捉长期依赖
  2. 预训练迁移 - ImageNet 权重提升泛化
  3. 实时推理 - 22ms 延迟,适合边缘部署

IMS 开发建议:

  • 采用 ViT-B/16 作为视觉骨干
  • 帧级特征聚合降低计算量
  • INT8 量化优化推理速度

本文基于 Nature Scientific Reports 2025 论文深度解读。


Transformer 实时疲劳检测:Nature 2025 论文深度解读与代码复现
https://dapalm.com/2026/08/16/2026-08-16-01-Transformer-Realtime-Fatigue-Detection/
作者
Mars
发布于
2026年8月16日
许可协议