TFormer:时频 Transformer 驾驶员疲劳识别的创新架构

TFormer:时频 Transformer 驾驶员疲劳识别的创新架构

一、论文信息

标题: TFormer: A time–frequency Transformer with batch normalization for driver fatigue recognition
期刊: Advanced Engineering Informatics (ScienceDirect), 2024
DOI: 10.1016/j.aei.2024.102234
发表时间: 2024年5月


二、时频分析的必要性

2.1 单一时域方法的局限

方法 局限性
时域 CNN 无法捕捉周期性特征
纯 Transformer 缺乏频率先验
手工时频特征 泛化能力差

2.2 时频域优势

  • 频率分辨率 - 识别眨眼频率、头部摆动周期
  • 抗噪声 - 频域滤波去除高频噪声
  • 多尺度分析 - 不同时间尺度的疲劳特征

三、TFormer 架构详解

3.1 整体框架

graph TB
    A[视频输入] --> B[关键点序列]
    B --> C1[时域分支]
    B --> C2[频域分支]
    
    C1 --> D1[时域 Transformer]
    C2 --> D2[频域 Transformer]
    
    D1 --> E[时频融合]
    D2 --> E
    
    E --> F[批归一化]
    F --> G[疲劳分类]

3.2 时频转换模块

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

class TimeFrequencyTransform(nn.Module):
"""
时域 → 频域转换

使用短时傅里叶变换(STFT)
"""

def __init__(self, n_fft=64, hop_length=16):
super().__init__()

self.n_fft = n_fft
self.hop_length = hop_length

def forward(self, x):
"""
Args:
x: 时域信号, shape=(B, T, D)

Returns:
freq_features: 频域特征, shape=(B, F, T', D)
"""
B, T, D = x.shape

# 对每个维度应用 STFT
freq_features = []

for d in range(D):
# 提取单维信号
signal = x[:, :, d] # (B, T)

# STFT
# 使用 torch.stft 实现
stft = torch.stft(
signal,
n_fft=self.n_fft,
hop_length=self.hop_length,
win_length=self.n_fft,
window=torch.hann_window(self.n_fft, device=x.device),
return_complex=True
)

# 取幅度谱
magnitude = torch.abs(stft) # (B, F, T')

freq_features.append(magnitude)

# 组合所有维度
freq_features = torch.stack(freq_features, dim=-1) # (B, F, T', D)

return freq_features


class TimeFrequencyTransformer(nn.Module):
"""
时频 Transformer 模块
"""

def __init__(self, embed_dim=256, num_heads=8, num_layers=4):
super().__init__()

# 1. 时域 Transformer
self.time_transformer = nn.TransformerEncoder(
nn.TransformerEncoderLayer(
d_model=embed_dim,
nhead=num_heads,
dim_feedforward=embed_dim * 4,
dropout=0.1,
batch_first=True
),
num_layers=num_layers
)

# 2. 频域 Transformer
self.freq_transformer = nn.TransformerEncoder(
nn.TransformerEncoderLayer(
d_model=embed_dim,
nhead=num_heads,
dim_feedforward=embed_dim * 4,
dropout=0.1,
batch_first=True
),
num_layers=num_layers
)

# 3. 时频融合
self.fusion = nn.Sequential(
nn.Linear(embed_dim * 2, embed_dim),
nn.ReLU(),
nn.Linear(embed_dim, embed_dim)
)

# 4. 批归一化(论文关键创新)
self.batch_norm = nn.BatchNorm1d(embed_dim)

# 5. 分类头
self.classifier = nn.Linear(embed_dim, 3)

def forward(self, time_features, freq_features):
"""
Args:
time_features: 时域特征, shape=(B, T, D)
freq_features: 频域特征, shape=(B, F, T', D)

Returns:
logits: 疲劳等级, shape=(B, 3)
"""
# 1. 时域编码
time_encoded = self.time_transformer(time_features) # (B, T, D)
time_pooled = time_encoded.mean(dim=1) # (B, D)

# 2. 频域编码
# Flatten 频率维度
B, F, T_prime, D = freq_features.shape
freq_flat = freq_features.view(B, F * T_prime, D)

freq_encoded = self.freq_transformer(freq_flat) # (B, F*T', D)
freq_pooled = freq_encoded.mean(dim=1) # (B, D)

# 3. 时频融合
fused = torch.cat([time_pooled, freq_pooled], dim=-1) # (B, 2D)
fused = self.fusion(fused) # (B, D)

# 4. 批归一化(关键创新)
# 提升训练稳定性
fused = fused.permute(0, 1) # (D, B) for BatchNorm1d
fused = self.batch_norm(fused)
fused = fused.permute(1, 0) # (B, D)

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

return logits


class TFormer(nn.Module):
"""
TFormer 完整模型
"""

def __init__(self, num_keypoints=68, embed_dim=256, num_frames=16):
super().__init__()

self.num_keypoints = num_keypoints
self.num_frames = num_frames

# 1. 关键点嵌入
self.keypoint_embed = nn.Linear(2, embed_dim)

# 2. 时频转换
self.tf_transform = TimeFrequencyTransform(n_fft=64, hop_length=16)

# 3. 时频 Transformer
self.tf_module = TimeFrequencyTransformer(
embed_dim=embed_dim,
num_heads=8,
num_layers=4
)

def forward(self, keypoint_seq):
"""
Args:
keypoint_seq: 关键点序列, shape=(B, T, N, 2)

Returns:
logits: 疲劳等级, shape=(B, 3)
"""
B, T, N, _ = keypoint_seq.shape

# 1. 关键点嵌入
x = self.keypoint_embed(keypoint_seq) # (B, T, N, D)

# 2. 时域特征(平均关键点)
time_features = x.mean(dim=2) # (B, T, D)

# 3. 频域特征
# 对每个关键点应用 STFT
freq_features = []
for n in range(N):
signal = keypoint_seq[:, :, n, :] # (B, T, 2)

# 归一化坐标到 [-1, 1]
signal_norm = (signal - 0.5) * 2

freq = self.tf_transform(signal_norm.view(B, T, 2))
freq_features.append(freq)

freq_features = torch.stack(freq_features, dim=2) # (B, F, T', N, D')
freq_features = freq_features.mean(dim=3) # (B, F, T', D')

# 4. 时频 Transformer
logits = self.tf_module(time_features, freq_features)

return logits


# 测试代码
if __name__ == "__main__":
model = TFormer(num_keypoints=68, embed_dim=256, num_frames=16)

# 模拟输入
keypoint_seq = torch.randn(2, 16, 68, 2)

# 前向传播
logits = model(keypoint_seq)

print(f"输入形状: {keypoint_seq.shape}")
print(f"输出形状: {logits.shape}")

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

四、批归一化的关键作用

4.1 为什么需要批归一化?

时频 Transformer 训练问题:

  • 时域和频域特征尺度差异大
  • 训练初期梯度不稳定
  • 收敛速度慢

4.2 批归一化的效果

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
def compare_training_curves():
"""
对比有无批归一化的训练曲线
"""
# 模拟训练曲线
epochs = np.arange(1, 101)

# 无批归一化:收敛慢,波动大
loss_no_bn = 1.0 / (1 + 0.05 * epochs) + 0.1 * np.random.randn(100)

# 有批归一化:收敛快,稳定
loss_with_bn = 1.0 / (1 + 0.08 * epochs) + 0.02 * np.random.randn(100)

print("无批归一化:")
print(f" 最终损失: {loss_no_bn[-1]:.4f}")
print(f" 收敛轮数: ~60")

print("\n有批归一化:")
print(f" 最终损失: {loss_with_bn[-1]:.4f}")
print(f" 收敛轮数: ~40")

五、实验结果

5.1 性能对比

方法 准确率 F1-score 推理时间
CNN-3D 89.2% 88.7% 35ms
ViT 91.5% 91.0% 28ms
Timeformer 93.8% 93.2% 25ms
TFormer 96.3% 95.9% 22ms

5.2 频域特征贡献

频段 疲劳特征 贡献度
0.1-0.5 Hz 头部摆动周期 25%
0.5-2.0 Hz 眨眼频率 35%
2.0-5.0 Hz 微表情 20%
>5.0 Hz 噪声 -

六、IMS 集成方案

6.1 系统架构

graph TB
    A[DMS 摄像头] --> B[关键点提取]
    B --> C[时域缓存]
    B --> D[频域分析]
    
    C --> E[时域 Transformer]
    D --> F[频域 Transformer]
    
    E --> G[时频融合]
    F --> G
    
    G --> H[批归一化]
    H --> I[疲劳分类]
    
    I --> J{疲劳等级}
    J -->|Level 0| K[正常]
    J -->|Level 1| L[警告]
    J -->|Level 2| M[严重]

6.2 开发检查清单

预处理:

  • 关键点提取(68 点)
  • 坐标归一化
  • 时域缓存(16 帧)

时频转换:

  • 实现 STFT
  • 提取幅度谱
  • 多频段特征

模型训练:

  • 训练时域 Transformer
  • 训练频域 Transformer
  • 批归一化配置

部署优化:

  • INT8 量化
  • 测试推理延迟
  • 验证批归一化效果

七、参考资源

  1. 论文原文: https://doi.org/10.1016/j.aei.2024.102234
  2. STFT 原理: https://en.wikipedia.org/wiki/Short-time_Fourier_transform
  3. Batch Normalization: https://arxiv.org/abs/1502.03167

八、总结

TFormer 实现96.3% 疲劳识别准确率,关键创新:

  1. 时频双域建模 - 捕捉时域和频域特征
  2. STFT 转换 - 提取频率域疲劳特征
  3. 批归一化 - 提升训练稳定性和收敛速度

IMS 开发建议:

  • 采用时频双分支架构
  • 重点提取 0.5-2.0 Hz 频段(眨眼)
  • 必须使用批归一化

本文基于 Advanced Engineering Informatics 2024 论文深度解读。


TFormer:时频 Transformer 驾驶员疲劳识别的创新架构
https://dapalm.com/2026/08/16/2026-08-16-04-TFormer-TimeFrequency-Transformer/
作者
Mars
发布于
2026年8月16日
许可协议