rPPG信号修复:WGAN生成式修复框架解决驾驶场景运动伪影——论文解读与代码复现

论文信息

项目 内容
标题 Inpainting for artifact restoration in rPPG signals: a tool for real-world applicability in wellness and driver monitoring systems
期刊 Physiological Measurement (IOP Science)
发表 2026年4月30日
链接 https://iopscience.iop.org/article/10.1088/1361-6579/ae60c1
PubMed https://pubmed.ncbi.nlm.nih.gov/41990813/
方法 WGAN + U-Net生成器
应用 PPG/rPPG信号退化段修复

核心创新

  1. 生成式修复rPPG信号退化段:首次将图像修复(inpainting)概念引入rPPG信号,用WGAN选择性重建被运动伪影损坏的信号段
  2. 大规模合成训练数据:生成覆盖广心率范围(40-180bpm)的合成PPG数据集,解决真实标注数据稀缺
  3. 保留生理波形结构:修复后信号不仅心率准确,还保留收缩峰、舒张波等波形形态学特征

问题定义

rPPG信号退化的三大场景

场景 退化原因 退化表现 发生频率
运动伪影 头部运动/车辆振动 信号完全淹没在噪声中 30-50%时间
光照突变 隧道出入口/树荫 信号基线漂移 10-20%时间
面部遮挡 手/物体遮住面部 信号缺失 5-15%时间

现有处理方式及局限

方法 原理 局限
线性插值 两端线性连接 完全忽略波形形态
样条插值 平滑曲线连接 无法恢复脉搏波形态
信号丢弃 丢弃退化段 心率监测中断
WGAN修复 生成式重建 需训练,但保留形态

方法详解

1. WGAN框架

flowchart LR
    A[退化rPPG信号] --> B[掩码标记退化段]
    B --> C[U-Net生成器G]
    C --> D[修复后信号]
    D --> E[判别器D]
    E --> F[Wasserstein损失]
    F --> G[梯度惩罚]
    G --> C

2. U-Net生成器

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
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
import torch
import torch.nn as nn
import torch.nn.functional as F

class PPGInpaintingGenerator(nn.Module):
"""
U-Net生成器:修复rPPG信号退化段

论文Section 3: Generator基于U-Net架构
输入:掩码后的退化信号
输出:修复后的完整信号
"""

def __init__(self, in_channels: int = 2, # 信号+掩码
out_channels: int = 1,
base_channels: int = 64):
super().__init__()

# 编码器(下采样)
self.enc1 = self._double_conv(in_channels, base_channels)
self.enc2 = self._double_conv(base_channels, base_channels * 2)
self.enc3 = self._double_conv(base_channels * 2, base_channels * 4)
self.enc4 = self._double_conv(base_channels * 4, base_channels * 8)

# 瓶颈层
self.bottleneck = self._double_conv(
base_channels * 8, base_channels * 16
)

# 解码器(上采样+跳连接)
self.upconv4 = nn.ConvTranspose1d(
base_channels * 16, base_channels * 8, 2, stride=2
)
self.dec4 = self._double_conv(
base_channels * 16, base_channels * 8
)

self.upconv3 = nn.ConvTranspose1d(
base_channels * 8, base_channels * 4, 2, stride=2
)
self.dec3 = self._double_conv(
base_channels * 8, base_channels * 4
)

self.upconv2 = nn.ConvTranspose1d(
base_channels * 4, base_channels * 2, 2, stride=2
)
self.dec2 = self._double_conv(
base_channels * 4, base_channels * 2
)

self.upconv1 = nn.ConvTranspose1d(
base_channels * 2, base_channels, 2, stride=2
)
self.dec1 = self._double_conv(
base_channels * 2, base_channels
)

# 输出层
self.out_conv = nn.Conv1d(base_channels, out_channels, 1)
self.out_act = nn.Tanh()

def _double_conv(self, in_ch, out_ch):
return nn.Sequential(
nn.Conv1d(in_ch, out_ch, 3, padding=1),
nn.BatchNorm1d(out_ch),
nn.LeakyReLU(0.2, inplace=True),
nn.Conv1d(out_ch, out_ch, 3, padding=1),
nn.BatchNorm1d(out_ch),
nn.LeakyReLU(0.2, inplace=True),
)

def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
"""
Args:
x: [B, 1, L] 退化rPPG信号
mask: [B, 1, L] 二值掩码(1=有效,0=退化段)

Returns:
restored: [B, 1, L] 修复后信号
"""
# 拼接信号和掩码
inp = torch.cat([x, mask], dim=1) # [B, 2, L]

# 编码器
e1 = self.enc1(inp) # [B, 64, L]
e2 = self.enc2(F.max_pool1d(e1, 2)) # [B, 128, L/2]
e3 = self.enc3(F.max_pool1d(e2, 2)) # [B, 256, L/4]
e4 = self.enc4(F.max_pool1d(e3, 2)) # [B, 512, L/8]

# 瓶颈
b = self.bottleneck(F.max_pool1d(e4, 2)) # [B, 1024, L/16]

# 解码器+跳连接
d4 = self.upconv4(b)
d4 = self.dec4(torch.cat([d4, e4], dim=1))

d3 = self.upconv3(d4)
d3 = self.dec3(torch.cat([d3, e3], dim=1))

d2 = self.upconv2(d3)
d2 = self.dec2(torch.cat([d2, e2], dim=1))

d1 = self.upconv1(d2)
d1 = self.dec1(torch.cat([d1, e1], dim=1))

# 输出
out = self.out_conv(d1)
out = self.out_act(out)

# 仅替换退化段,保留有效段原始信号
restored = x * mask + out * (1 - mask)

return restored


class WGANCritic(nn.Module):
"""
WGAN判别器(Critic)

使用Wasserstein距离而非JS散度,
梯度惩罚确保Lipschitz连续性
"""

def __init__(self, in_channels: int = 1):
super().__init__()
self.model = nn.Sequential(
nn.Conv1d(in_channels, 64, 3, 2, 1),
nn.LeakyReLU(0.2),

nn.Conv1d(64, 128, 3, 2, 1),
nn.BatchNorm1d(128),
nn.LeakyReLU(0.2),

nn.Conv1d(128, 256, 3, 2, 1),
nn.BatchNorm1d(256),
nn.LeakyReLU(0.2),

nn.Conv1d(256, 512, 3, 2, 1),
nn.BatchNorm1d(512),
nn.LeakyReLU(0.2),

nn.AdaptiveAvgPool1d(1),
nn.Flatten(),
nn.Linear(512, 1),
)

def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.model(x)


def gradient_penalty(critic, real, fake, mask, device='cpu'):
"""
梯度惩罚(WGAN-GP)

确保判别器满足Lipschitz连续性约束
"""
batch_size = real.size(0)

# 随机插值因子
alpha = torch.rand(batch_size, 1, 1, device=device)

# 仅在退化段计算梯度惩罚
interpolated = (alpha * real + (1 - alpha) * fake).requires_grad_(True)

critic_input = interpolated * (1 - mask) + real * mask
critic_output = critic(critic_input)

gradients = torch.autograd.grad(
outputs=critic_output,
inputs=interpolated,
grad_outputs=torch.ones_like(critic_output),
create_graph=True,
retain_graph=True,
)[0]

gradients = gradients.view(batch_size, -1)
gradient_norm = gradients.norm(2, dim=1)
penalty = ((gradient_norm - 1) ** 2).mean()

return penalty


# 训练循环
def train_wgan_inpainting(
generator, critic,
dataloader, epochs=100,
lr_g=1e-4, lr_c=1e-4,
lambda_gp=10,
device='cpu'
):
"""
WGAN修复模型训练

论文Section 3.3:
- Generator: Adam(lr=1e-4, beta1=0.5, beta2=0.9)
- Critic: Adam(lr=1e-4, beta1=0.5, beta2=0.9)
- 梯度惩罚系数: 10
- Critic更新次数/Generator: 5
"""
opt_g = torch.optim.Adam(
generator.parameters(), lr=lr_g, betas=(0.5, 0.9)
)
opt_c = torch.optim.Adam(
critic.parameters(), lr=lr_c, betas=(0.5, 0.9)
)

for epoch in range(epochs):
for batch in dataloader:
clean, degraded, mask = batch
clean = clean.to(device)
degraded = degraded.to(device)
mask = mask.to(device)

# 1. 训练Critic(5次)
for _ in range(5):
opt_c.zero_grad()

fake = generator(degraded, mask).detach()

# Wasserstein距离
real_score = critic(clean * (1 - mask) + degraded * mask)
fake_score = critic(fake * (1 - mask) + degraded * mask)

gp = gradient_penalty(
critic, clean, fake, mask, device
)

loss_c = -(real_score.mean() - fake_score.mean()) + lambda_gp * gp
loss_c.backward()
opt_c.step()

# 2. 训练Generator
opt_g.zero_grad()

fake = generator(degraded, mask)
fake_score = critic(fake * (1 - mask) + degraded * mask)

# 对抗损失 + 重建损失
loss_adv = -fake_score.mean()
loss_recon = F.l1_loss(
fake * (1 - mask), clean * (1 - mask)
)

loss_g = loss_adv + 100 * loss_recon # 重建权重
loss_g.backward()
opt_g.step()

if (epoch + 1) % 10 == 0:
print(f"Epoch {epoch+1}: G_loss={loss_g:.4f}, "
f"C_loss={loss_c:.4f}, "
f"recon={loss_recon:.4f}")


# 合成数据生成
def generate_synthetic_ppg(
n_samples=10000, length=1800, fs=30
):
"""
生成合成PPG信号用于训练

论文Section 3.2: 大规模合成数据集
心率范围: 40-180 bpm
包含: 基线漂移、呼吸调制、二阶波
"""
np.random.seed(42)
signals = []
masks = []

for i in range(n_samples):
# 随机心率
hr = np.random.uniform(40, 180)
freq = hr / 60 # Hz

# 时间轴
t = np.arange(length) / fs

# 基础脉搏波(高斯导数模型)
pulse = np.exp(-((t * freq % 1) - 0.3) ** 2 / 0.01)
pulse -= 0.5 * np.exp(-((t * freq % 1) - 0.5) ** 2 / 0.02)

# 呼吸调制
resp = np.sin(2 * np.pi * 0.25 * t) * 0.1

# 基线漂移
baseline = np.cumsum(np.random.randn(length)) * 0.001

# 噪声
noise = np.random.randn(length) * 0.05

signal = pulse + resp + baseline + noise

# 随机生成退化段
mask = np.ones(length)
n_artifacts = np.random.randint(1, 5)
for _ in range(n_artifacts):
start = np.random.randint(0, length - 100)
duration = np.random.randint(30, 200)
mask[start:start+duration] = 0
# 在退化段添加大噪声
signal[start:start+duration] += np.random.randn(duration) * 2.0

signals.append(signal)
masks.append(mask)

return np.array(signals, dtype=np.float32), \
np.array(masks, dtype=np.float32)


if __name__ == "__main__":
import numpy as np

# 生成合成数据
print("生成合成PPG数据...")
signals, masks = generate_synthetic_ppg(n_samples=100, length=600)

# 创建退化信号
degraded = signals * masks + np.random.randn(*signals.shape) * 0.5 * (1 - masks)

# 初始化模型
generator = PPGInpaintingGenerator()
critic = WGANCritic()

print(f"生成器参数: {sum(p.numel() for p in generator.parameters()):,}")
print(f"判别器参数: {sum(p.numel() for p in critic.parameters()):,}")

# 前向测试
x = torch.from_numpy(degraded[:2]).unsqueeze(1) # [2, 1, 600]
m = torch.from_numpy(masks[:2]).unsqueeze(1)

restored = generator(x, m)
print(f"输入: {x.shape}")
print(f"掩码: {m.shape}")
print(f"修复: {restored.shape}")
print(f"退化段MAE: {F.l1_loss(restored[~m.bool()], x[~m.bool()]):.4f}")

实验结果

修复质量评估

退化类型 修复前MAE(bpm) 修复后MAE(bpm) 改善
运动伪影 22.5 5.8 74%
光照突变 15.3 3.2 79%
面部遮挡 18.7 4.5 76%
混合退化 19.8 5.1 74%

波形形态学保持

指标 线性插值 样条插值 WGAN修复
心率MAE 8.2 6.5 5.1
波形相似度 0.45 0.58 0.82
收缩峰检测率 35% 48% 78%
舒张波检测率 22% 35% 65%

IMS开发启示

1. rPPG信号质量评估管道

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
class rPPGQualityAssessment:
"""
rPPG信号质量评估+修复管道

1. 实时检测信号退化段
2. WGAN修复退化段
3. 输出连续心率估计
"""

def __init__(self, threshold=0.3):
self.quality_threshold = threshold
self.inpaintor = PPGInpaintingGenerator()
# 加载预训练权重
# self.inpaintor.load_state_dict(...)

def assess_quality(self, rppg_signal):
"""
评估rPPG信号质量

返回二值掩码:1=有效, 0=退化
"""
# 基于SNR的退化检测
window = 30 # 1秒@30fps
mask = np.ones(len(rppg_signal))

for i in range(0, len(rppg_signal) - window, window//2):
segment = rppg_signal[i:i+window]
snr = self._compute_snr(segment)
if snr < self.quality_threshold:
mask[i:i+window] = 0

return mask

def _compute_snr(self, segment):
"""计算信噪比"""
# FFT找到心率峰值
fft = np.abs(np.fft.rfft(segment))
freqs = np.fft.rfftfreq(len(segment), d=1/30)

# 心率频段(0.7-4Hz)
hr_mask = (freqs >= 0.7) & (freqs <= 4.0)
signal_power = np.max(fft[hr_mask])
noise_power = np.median(fft[~hr_mask])

return signal_power / (noise_power + 1e-8)

def restore(self, rppg_signal):
"""修复退化信号"""
mask = self.assess_quality(rppg_signal)

if (mask == 0).sum() == 0:
return rppg_signal # 无退化

# WGAN修复
x = torch.from_numpy(rppg_signal).float().unsqueeze(0).unsqueeze(0)
m = torch.from_numpy(mask).float().unsqueeze(0).unsqueeze(0)

with torch.no_grad():
restored = self.inpaintor(x, m)

return restored.squeeze().numpy()

2. 与MS-rPPG的协同

组件 功能 前一篇论文 本篇论文
多光谱融合 RGB+NIR ✅ MS-rPPG -
信号修复 退化段重建 - ✅ WGAN
心率估计 信号→心率 ✅ MS-Mamba -
联合管道 端到端 MS-rPPG+本方法 -

3. 部署方案

阶段 方案 延迟 精度
P0 丢弃退化段 高(断续)
P1 线性插值
P2 WGAN修复 50ms
P3 MS-rPPG+WGAN 100ms 最高

总结

rPPG信号修复是驾驶场景心率监测的关键使能技术:

  1. WGAN修复MAE 5.1 bpm:比线性插值(8.2)和样条(6.5)分别改善38%和22%
  2. 保留波形形态:收缩峰检测率78%,线性插值仅35%
  3. 合成数据训练:无需大量真实标注数据
  4. 与MS-rPPG协同:先修复退化段,再做多光谱融合,实现全天候连续心率监测

https://dapalm.com/2026/09/21/2026-09-21-19-rppg-wgan-inpainting-artifact-restoration-ims/
作者
Mars
发布于
2026年9月21日
许可协议