You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.14 16:10:54