基于UNet的键盘检测算法真实场景性能不佳的优化咨询
键盘按键与主体检测算法:合成数据适配真实场景的解决方案
问题背景
我正在开发一款针对已知型号键盘的按键与键盘主体检测算法,通过Blender搭建环境生成带随机光照、角度、纹理及屏幕遮挡物的图像作为训练数据。采用自定义结构的UNet模型(代码如下),基于500张生成图像训练后,模型在合成训练/测试数据上表现优异,但在真实照片中性能极差。已尝试调参、调暗Blender生成图像、用OpenCV处理结果,效果均不理想。想请教除优化Blender环境外的解决办法,以及是否可更换神经网络类型或采用其他方案。
自定义UNet模型代码
import torch import torch.nn as nn class RELUConvBlock(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() layers = [ nn.Conv2d(in_ch, out_ch,3,1,1), nn.BatchNorm2d(out_ch), nn.ReLU() ] self.model = nn.Sequential(*layers) def forward(self,x): return self.model(x) class DownBlock(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() layers = [ RELUConvBlock(in_ch, out_ch), RELUConvBlock(out_ch, out_ch) ] self.model = nn.Sequential(*layers) def forward(self,x): return self.model(x) class UpBlock(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() layers = [ RELUConvBlock(in_ch, out_ch), RELUConvBlock(out_ch, out_ch) ] self.model = nn.Sequential(*layers) def forward(self,x): return self.model(x) class UNet(nn.Module): def __init__(self, out_ch = 3, down_ch = [64,128,256,512]): super().__init__() self.pool = nn.MaxPool2d(kernel_size=(2,2), stride=(2,2)) self.down0 = DownBlock(3, down_ch[0]) self.down1 = DownBlock(down_ch[0], down_ch[1]) self.down2 = DownBlock(down_ch[1], down_ch[2]) self.down3 = DownBlock(down_ch[2], down_ch[3]) self.bottleneck = DownBlock(down_ch[3], 2*down_ch[3]) self.up3 = UpBlock(2*down_ch[-1], down_ch[-1]) self.up2 = UpBlock(down_ch[-1], down_ch[-2]) self.up1 = UpBlock(down_ch[-2], down_ch[-3]) self.up0 = UpBlock(down_ch[-3], down_ch[-4]) self.connect_b_up3 = nn.ConvTranspose2d(2*down_ch[-1], down_ch[-1],kernel_size=2,stride=2) self.connect_up3_up2 = nn.ConvTranspose2d(down_ch[-1], down_ch[-2],kernel_size=2,stride=2) self.connect_up2_up1 = nn.ConvTranspose2d(down_ch[-2], down_ch[-3],kernel_size=2,stride=2) self.connect_up1_up0 = nn.ConvTranspose2d(down_ch[-3], down_ch[-4],kernel_size=2,stride=2) self.final_conv = nn.Conv2d(down_ch[0], out_ch, kernel_size=1) def _crop_to_match(self, tensor, target): _, _, h, w = target.shape return tensor[:, :, :h, :w] def forward(self, x): skip_connections = [] x = self.down0(x) skip_connections.append(x) x = self.pool(x) x = self.down1(x) skip_connections.append(x) x = self.pool(x) x = self.down2(x) skip_connections.append(x) x = self.pool(x) x = self.down3(x) skip_connections.append(x) x = self.pool(x) x = self.bottleneck(x) x = self.connect_b_up3(x) skip_connection = self._crop_to_match(skip_connections[3], x) x = torch.cat((skip_connection, x), dim=1) x = self.up3(x) x = self.connect_up3_up2(x) skip_connection = self._crop_to_match(skip_connections[2], x) x = torch.cat((skip_connection, x), dim=1) x = self.up2(x) x = self.connect_up2_up1(x) skip_connection = self._crop_to_match(skip_connections[1], x) x = torch.cat((skip_connection, x), dim=1) x = self.up1(x) x = self.connect_up1_up0(x) skip_connection = self._crop_to_match(skip_connections[0], x) x = torch.cat((skip_connection, x), dim=1) x = self.up0(x) x = self.final_conv(x) return x #nn.functional.sigmoid(x)
训练代码
from torch.optim import Adam from tqdm import tqdm LEARNING_RATE = 1e-4 num_epochs = 10 loss_fn = nn.CrossEntropyLoss() optimizer = Adam(model.parameters(), lr=LEARNING_RATE) scaler = torch.amp.GradScaler(device) model.train() for epoch in range(num_epochs): loop = tqdm(enumerate(train_loader), total = len(train_loader)) for batch_idx, (data, targets) in loop: data = data.to(device) targets = targets.to(device) with torch.amp.autocast(device_type=str(device)): predictions = model(data) loss = loss_fn(predictions, targets) optimizer.zero_grad() scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() loop.set_postfix(loss = loss.item())
解决方案
一、数据适配策略
- 少量真实数据微调:标注10-50张真实键盘照片,用合成数据训练好的UNet做迁移学习。冻结模型大部分底层特征提取层,只训练顶层分割层或瓶颈层,让模型快速适应真实场景的纹理、光照噪声。
- 增强合成数据的真实感:对合成数据做更贴近真实拍摄的增强操作:
- 添加随机高斯噪声、椒盐噪声,模拟照片传感器噪点
- 随机调整色温、对比度、饱和度,匹配不同真实拍摄环境
- 加入高斯模糊、运动模糊,模拟手持拍摄的模糊效果
- 叠加真实背景(如桌面、书本、办公场景),替换Blender的纯色背景
- 无监督领域自适应:无需标注真实数据,用领域对抗神经网络(DANN)让模型学习跨域不变特征,降低合成数据与真实数据的分布差异;或用CycleGAN将合成图像风格转换为真实照片风格,再用转换后的图像训练模型。
二、模型结构优化
- 更换分割模型:
- Attention UNet/UNet++:在基础UNet上加入注意力机制或嵌套结构,提升对按键这类小目标的检测能力,增强特征复用效率
- Mask R-CNN:如果需要同时检测键盘主体和每个按键的实例,这款实例分割模型更合适,它对真实场景的鲁棒性更强,能输出精确的目标掩码
- SegNet:采用编码器-解码器结构,通过池化索引进行上采样,减少特征丢失,对按键边缘细节的保留效果更好
- 引入预训练骨干网络:将自定义UNet的编码器替换为在海量真实图像上预训练的ResNet、EfficientNet等模型,这类预训练模型自带真实场景的特征提取能力,能大幅提升模型对真实数据的适配性。
三、训练策略调整
- 优化损失函数:将单一的CrossEntropyLoss替换为Dice Loss + CrossEntropyLoss组合,Dice Loss更关注掩码的重叠度,适合精细分割任务;如果是多分类场景,可尝试Focal Loss,降低大面积背景等简单样本的权重,让模型聚焦于难分割的按键区域。
- 调整训练轮次与学习率:当前仅训练10轮,可适当增加到20-50轮,并加入学习率调度(如CosineAnnealingLR、StepLR),让模型在训练后期稳定收敛;微调真实数据时,使用更小的学习率(如1e-5)。
- 半监督训练:若真实数据标注成本高,可采用半监督策略:用合成数据训练的模型对未标注真实数据生成伪标签,将真实数据+伪标签与合成数据混合训练,逐步提升模型对真实场景的适应能力。
四、后处理优化
- 除基础OpenCV操作外,加入针对性后处理:
- 用形态学膨胀/腐蚀消除分割后的小噪点
- 通过轮廓检测筛选符合键盘形状的区域,过滤误分割部分
- 用连通域分析保留面积符合按键大小的区域,剔除异常小区域
内容的提问来源于stack exchange,提问作者Furnitcher
相关产品推荐
相关产品推荐

