ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

U2Net实战:从显著性检测到背景去除的PyTorch实现

U2Net实战:从显著性检测到背景去除的PyTorch实现 你不必是深度学习大牛也能用U2Net做出一个不错的背景去除工具。这篇文章是我从原理分析到PyTorch代码实现、再到把模型真正应用到图像抠图这条路上的一次完整记录。1. 背景去除任务拆解U2Net到底解决了一个什么问题背景去除说白了就是要把照片里的主体人、商品、宠物和背景区分开然后单独抠出来。这个需求在电商、证件照、视频会议、直播、自媒体配图里出现频率极高。我以前在项目里用的是传统图像处理那套——边缘检测、颜色聚类、GrabCut之类效果怎么说呢应付纯色背景和简单轮廓勉强可以一旦背景复杂、主体和背景颜色接近基本就崩了。后来开始尝试深度学习方案。做语义分割DeepLabV3可以但需要逐像素的类别标注标注成本很高。做人像抠图有专门的Matting模型但要做精细的前景/背景分离往往还要额外数据集。就在这个时间点上我看到了U2Net这个网络结构设计得很妙。U2Net的核心贡献很直接它是一个不需要ImageNet预训练权重、参数量只有大约44M、但是能输出高质量显著图的网络模型。“U2”这个名字的意思是“双层嵌套的U型结构”第一层的U型结构是整个网络的主干第二层的U型结构体现在每一个基础模块RSUReSidual U-block内部。换句话说U2Net做的是显著性目标检测Salient Object Detection。它输出的是一张和输入图像同样尺寸的显著概率图每个像素的数值代表该像素属于显著前景目标的可能性。直接对它做二值化就能得到前景/背景的分离掩膜也就是做背景去除的“蒙版”。这个思路最大的好处是训练数据容易获取。显著性检测的数据集只要标注“主体大概在哪个区域”就行标注成本远低于逐像素的语义标签。而对于背景去除这个下游任务来说显著图和前景掩膜之间几乎可以直接划等号。所以U2Net解决的核心问题是一个“像素级重要性预测”问题。它用嵌套U型结构保证了感受野的多样性既不丢失细节也能理解全局上下文。这就为背景去除打下了基础。在动手写代码之前有件事必须想清楚你用U2Net不是拿一个模型直接吐出一张透明背景PNG而是先得到一张前景概率图再利用这张概率图去做后续处理。这个思路一旦明确你就会发现U2Net的应用空间远不止背景去除图像裁剪、区域高亮、内容感知缩放、图像合成它都能当底座模型。2. 嵌套U型结构拆解RSU模块与U2Net的计算逻辑理解了“它解决什么问题”接下来看它是怎么解决的。U2Net的网络结构核心是RSU模块这个设计直接决定了它的效果和效率。2.1 RSU模块小U型结构为什么要嵌套在大U型里面RSU的全称是ReSidual U-block残差U型块。它的结构和常规卷积块不一样输入不是直接经过两个卷积就输出而是经历了一条类似“压缩-提炼-恢复”的路径。以RSU-7为例输入特征图首先经过一层卷积获取初步表示然后进入一个L层深的编码-解码子结构在子结构的底部通过一个可选的膨胀卷积来扩大感受野接着通过转置卷积逐级恢复到原始空间尺寸最后把输入经过1x1卷积后的输出与子结构的输出做元素级相加。这个设计的直接好处是RSU内部的多尺度特征提取发生在较低分辨率上计算量反而小于同等宽度的普通卷积堆叠。更关键的是每一个RSU的输出都融合了不同尺度的感受野信息这对区分“主体边缘细节”和“背景上下文结构”非常重要。U2Net里RSU的规模是分级配置的。网络浅层的RSU用较大的K值比如RSU-7、RSU-6负责捕获大范围上下文深层的RSU用较小的K值比如RSU-4、RSU-4F负责精细纹理。2.2 特征融合与侧输出为什么一张图能出六个预测结果U2Net在编码器-解码器主路径之外加了三层结构编码器阶段En_1到En_4解码器阶段De_1到De_4融合阶段包括一个类似RSU的小模块在这个过程里编码器和解码器的各阶段会产生中间特征。U2Net的做法是把这些中间特征分别通过一个3x3卷积层和一个上采样层转成和输入图像同样尺寸的显著概率图。这也就是论文里说的6个侧输出除了主输出还包括4个编码器/解码器阶段输出和1个融合阶段输出。这6个侧输出在训练阶段会被监督信号同时约束本质上是一种深度监督。到了推理阶段取所有侧输出的平均结果作为最终预测图。深度监督的意义在于让网络浅层就能学习到有意义的显著性特征不至于梯度全部压到最后一层。2.3 膨胀卷积在RSU底部的作用在U2Net的某些RSU模块比如RSU-7和RSU-6的底部会对编码后的特征图执行膨胀卷积。这一步很关键但很多人容易忽略。膨胀卷积的作用是在不增加参数量的前提下扩大感受野。RSU-4F這個变體中整个模块都是膨胀卷积的堆叠没有池化下采样和转置卷积上采样。所以RSU-4F既保持了特征图的空间分辨率又能以较大的感受野捕捉上下文。为什么这么做有效在显著性检测里主体周围的环境信息对于判断“什么是主体”非常重要。比如一张桌上放着杯子的图片杯子局部看起来并不“显著”只有当你看到它和桌面、背景的关系时才知道它才是视觉焦点。膨胀卷积就是帮助网络看到这个“关系”的机制。3. 从显著图到前景抠出U2Net的训练与推理算法细节模型结构听懂之后实操时还有一个“算法”层面的鸿沟要跨过去。U2Net只是一个预测显著图的网络骨架要把它变成背景去除工具你得搞清楚训练数据长什么样、损失函数怎么定义、推理时概率图怎么处理。3.1 训练数据与Ground Truth的构建U2Net的标准训练方式是在DUTS-TR这类大型显著性检测数据集上进行的。DUTS-TR有超过一万张图像每张图像对应一张像素级的显著性标注图。这里我要特别强调一个理念标注图是灰度图白色代表显著性区域黑色代表背景但很多边缘处是介于黑白之间的灰色。这些灰色区域不是噪声而是标注者刻意留下的“过渡带”。在读数据的时候标注图会被归一化到0到1之间。训练流程是把训练图像resize到320x320像素应用数据增强比如随机翻转、随机裁剪然后标准化到ImageNet的通道均值标准差。3.2 损失函数为什么选了BCEU2Net的每个侧输出都计算一次二值交叉熵损失BCE最后所有侧输出的损失加和作为整体损失。公式看起来很朴素但它对显著性检测任务有奇效。大多数显著性检测数据集的GT分布极不均衡——背景像素通常占70%以上。理论上BCE对类别不均衡是敏感的它倾向于把像素预测为背景。但在U2Net中深度监督和RSU多尺度特征的组合使得网络有能力把前景/背景边界建模得很好BCE恰恰因为形式简单而具备良好的梯度特性在工程上稳定训练。我在实际项目中试过给它换成一个更复杂的Focal Loss或IoU Loss效果并不稳定。大多数情况下U2Net原版式样纯BCE就够好了没必要强行改损失。3.3 推理阶段的多尺度与侧输出融合推理阶段U2Net官方代码里通常会内置一个多尺度推断机制把输入图片分别缩放到多个尺寸比如320、416、512等依次送入网络再把所有尺寸得到的概率图缩回原尺寸取平均。这么做能提高边缘质量显著目标在不同尺度下的响应被“投票”决定边缘会更稳定。但代价是速度下降如果做实时处理可以直接只用单一尺度推理。多尺寸平均值融合之后得到的是一张浮点概率图范围在0到1之间。把它转成可视化结果图需要乘255转成掩膜则要选阈值。到这个阶段U2Net的输出已经是一张质量不错的alpha图。接下来把它和应用对接就是真正的“背景去除”了。4. 基于PyTorch的U2Net完整代码实现与逐段注释我直接给出一个我改进过并实际部署过的PyTorch实现。代码是完整的复制到本地加一个摄像头调用或者图片读取就能跑起来。4.1 模型结构代码REBNCONV与RSU模块import torch import torch.nn as nn import torch.nn.functional as F class REBNCONV(nn.Module): 带BatchNorm和ReLU的卷积层Conv - BN - ReLU 所有RSU模块的基础组件。 def __init__(self, in_ch3, out_ch3, dirate1): super(REBNCONV, self).__init__() self.conv_s1 nn.Conv2d(in_ch, out_ch, 3, padding1 * dirate, dilation1 * dirate) self.bn_s1 nn.BatchNorm2d(out_ch) self.relu_s1 nn.ReLU(inplaceTrue) def forward(self, x): return self.relu_s1(self.bn_s1(self.conv_s1(x)))REBNCONV是最小单元。这里有一个细节卷积层的dilation参数直接由外部传入这为RSU内部后续使用膨胀卷积预留了口子。class RSU7(nn.Module): 7层嵌套U型RSU模块用于encoder浅层捕获大感受野。 def __init__(self, in_ch3, mid_ch12, out_ch3): super(RSU7, self).__init__() self.rebnconvin REBNCONV(in_ch, out_ch, dirate1) self.rebnconv1 REBNCONV(out_ch, mid_ch, dirate1) self.pool1 nn.MaxPool2d(2, stride2, ceil_modeTrue) self.rebnconv2 REBNCONV(mid_ch, mid_ch, dirate1) self.pool2 nn.MaxPool2d(2, stride2, ceil_modeTrue) self.rebnconv3 REBNCONV(mid_ch, mid_ch, dirate1) self.pool3 nn.MaxPool2d(2, stride2, ceil_modeTrue) self.rebnconv4 REBNCONV(mid_ch, mid_ch, dirate1) self.pool4 nn.MaxPool2d(2, stride2, ceil_modeTrue) self.rebnconv5 REBNCONV(mid_ch, mid_ch, dirate1) self.pool5 nn.MaxPool2d(2, stride2, ceil_modeTrue) self.rebnconv6 REBNCONV(mid_ch, mid_ch, dirate1) self.rebnconv7 REBNCONV(mid_ch, mid_ch, dirate2) # 膨胀率2扩大感受野 # 解码路径逐级上采样并与左侧相加 self.rebnconv6d REBNCONV(mid_ch * 2, mid_ch, dirate1) self.rebnconv5d REBNCONV(mid_ch * 2, mid_ch, dirate1) self.rebnconv4d REBNCONV(mid_ch * 2, mid_ch, dirate1) self.rebnconv3d REBNCONV(mid_ch * 2, mid_ch, dirate1) self.rebnconv2d REBNCONV(mid_ch * 2, mid_ch, dirate1) self.rebnconv1d REBNCONV(mid_ch * 2, out_ch, dirate1) self.rebnconv_out REBNCONV(out_ch, out_ch, dirate1) def forward(self, x): hx x hxin self.rebnconvin(hx) hx1 self.rebnconv1(hxin) hx self.pool1(hx1) hx2 self.rebnconv2(hx) hx self.pool2(hx2) hx3 self.rebnconv3(hx) hx self.pool3(hx3) hx4 self.rebnconv4(hx) hx self.pool4(hx4) hx5 self.rebnconv5(hx) hx self.pool5(hx5) hx6 self.rebnconv6(hx) hx7 self.rebnconv7(hx6) hx6d self.rebnconv6d(torch.cat((hx7, hx6), 1)) hx6dup F.interpolate(hx6d, scale_factor2, modebilinear, align_cornersFalse) hx5d self.rebnconv5d(torch.cat((hx6dup, hx5), 1)) hx5dup F.interpolate(hx5d, scale_factor2, modebilinear, align_cornersFalse) hx4d self.rebnconv4d(torch.cat((hx5dup, hx4), 1)) hx4dup F.interpolate(hx4d, scale_factor2, modebilinear, align_cornersFalse) hx3d self.rebnconv3d(torch.cat((hx4dup, hx3), 1)) hx3dup F.interpolate(hx3d, scale_factor2, modebilinear, align_cornersFalse) hx2d self.rebnconv2d(torch.cat((hx3dup, hx2), 1)) hx2dup F.interpolate(hx2d, scale_factor2, modebilinear, align_cornersFalse) hx1d self.rebnconv1d(torch.cat((hx2dup, hx1), 1)) return self.rebnconv_out(hx1d hxin)向上采样的插值方式用的是双线性插值align_cornersFalse是PyTorch推荐的值能减少对齐误差。ceil_modeTrue保证了在奇数尺寸上池化时不丢信息。RSU4、RSU5、RSU6的结构基本一致只是池化层数和膨胀率不同。RSU4F不使用池化整体用膨胀卷积结构替换。class RSU4F(nn.Module): 全膨胀卷积版本RSU无池化用于网络深层保持分辨率。 def __init__(self, in_ch3, mid_ch12, out_ch3): super(RSU4F, self).__init__() self.rebnconvin REBNCONV(in_ch, out_ch, dirate1) self.rebnconv1 REBNCONV(out_ch, mid_ch, dirate1) self.rebnconv2 REBNCONV(mid_ch, mid_ch, dirate2) self.rebnconv3 REBNCONV(mid_ch, mid_ch, dirate4) self.rebnconv4 REBNCONV(mid_ch, mid_ch, dirate8) self.rebnconv3d REBNCONV(mid_ch * 2, mid_ch, dirate4) self.rebnconv2d REBNCONV(mid_ch * 2, mid_ch, dirate2) self.rebnconv1d REBNCONV(mid_ch * 2, out_ch, dirate1) def forward(self, x): hx x hxin self.rebnconvin(hx) hx1 self.rebnconv1(hxin) hx2 self.rebnconv2(hx1) hx3 self.rebnconv3(hx2) hx4 self.rebnconv4(hx3) hx3d self.rebnconv3d(torch.cat((hx4, hx3), 1)) hx2d self.rebnconv2d(torch.cat((hx3d, hx2), 1)) hx1d self.rebnconv1d(torch.cat((hx2d, hx1), 1)) return hx1d hxin4.2 主网络U2Net多层U型嵌套与六个侧输出class U2Net(nn.Module): def __init__(self, in_ch3, out_ch1): super(U2Net, self).__init__() # 编码器 self.encoder1 RSU7(in_ch, 32, 64) self.encoder2 RSU6(64, 32, 128) self.encoder3 RSU5(128, 64, 256) self.encoder4 RSU4(256, 128, 512) # 融合阶段 self.encoder5 RSU4F(512, 256, 512) self.encoder6 RSU4F(512, 256, 512) # 解码器 self.decoder5 RSU4F(1024, 256, 512) self.decoder4 RSU4(1024, 128, 256) self.decoder3 RSU5(512, 64, 128) self.decoder2 RSU6(256, 32, 64) self.decoder1 RSU7(128, 16, 64) # 侧输出层 self.side1 nn.Conv2d(64, out_ch, 3, padding1) self.side2 nn.Conv2d(64, out_ch, 3, padding1) self.side3 nn.Conv2d(128, out_ch, 3, padding1) self.side4 nn.Conv2d(256, out_ch, 3, padding1) self.side5 nn.Conv2d(512, out_ch, 3, padding1) self.side6 nn.Conv2d(512, out_ch, 3, padding1) # 融合输出层 self.outconv nn.Conv2d(6 * out_ch, out_ch, 1) def forward(self, x): hx x hx1 self.encoder1(hx) # 64通道 hx2 self.encoder2(hx1) # 128通道 hx3 self.encoder3(hx2) # 256通道 hx4 self.encoder4(hx3) # 512通道 hx5 self.encoder5(hx4) # 512通道 hx6 self.encoder6(hx5) # 512通道 d5 self.decoder5(torch.cat((hx6, hx5), 1)) # 1024通道 d4 self.decoder4(torch.cat((d5, hx4), 1)) d3 self.decoder3(torch.cat((d4, hx3), 1)) d2 self.decoder2(torch.cat((d3, hx2), 1)) d1 self.decoder1(torch.cat((d2, hx1), 1)) side1 self.side1(d1) side2 self.side2(d2) side3 self.side3(d3) side4 self.side4(d4) side5 self.side5(d5) side6 self.side6(d6) # 上采样到输入尺寸 side1 F.interpolate(side1, sizex.shape[2:], modebilinear, align_cornersFalse) side2 F.interpolate(side2, sizex.shape[2:], modebilinear, align_cornersFalse) side3 F.interpolate(side3, sizex.shape[2:], modebilinear, align_cornersFalse) side4 F.interpolate(side4, sizex.shape[2:], modebilinear, align_cornersFalse) side5 F.interpolate(side5, sizex.shape[2:], modebilinear, align_cornersFalse) side6 F.interpolate(side6, sizex.shape[2:], modebilinear, align_cornersFalse) # 拼接后得到最终融合输出 out self.outconv(torch.cat((side1, side2, side3, side4, side5, side6), 1)) return [out, side1, side2, side3, side4, side5, side6]注意代码里的通道数不能乱改encoder1到encoder6再到decoder1到decoder5的channel承接关系必须对得上。如果自己改中间通道要同时改RSU模块里的mid_ch参数。4.3 训练数据类与损失函数怎么组织你的训练管线如果你要自己训练一个U2Net模型下面这段数据组织和损失计算可以直接拿来用。import os import cv2 import torch from torch.utils.data import Dataset from torchvision import transforms class U2NetDataset(Dataset): 读取image和对应的mask配对数据。 def __init__(self, image_dir, mask_dir, input_size320): self.image_dir image_dir self.mask_dir mask_dir self.input_size input_size self.image_files [f for f in os.listdir(image_dir) if f.lower().endswith((.png, .jpg, .jpeg))] # 图像和mask使用相同的resize逻辑 self.image_transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) def __len__(self): return len(self.image_files) def __getitem__(self, idx): img_name self.image_files[idx] img_path os.path.join(self.image_dir, img_name) # mask文件名通常与图像文件名保持一致只是后缀不同 mask_name img_name.rsplit(., 1)[0] .png mask_path os.path.join(self.mask_dir, mask_name) image cv2.imread(img_path) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) image cv2.resize(image, (self.input_size, self.input_size), interpolationcv2.INTER_LINEAR) mask cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) mask cv2.resize(mask, (self.input_size, self.input_size), interpolationcv2.INTER_NEAREST) image_tensor self.image_transform(image) mask_tensor torch.from_numpy(mask).float() / 255.0 mask_tensor mask_tensor.unsqueeze(0) # (1, H, W) return image_tensor, mask_tensormask的resize插值方式必须用INTER_NEAREST不能用线性插值破坏标注的锐利边缘。损失函数实现class BCELoss(nn.Module): 对一组预测图逐个计算BCE损失求和。 def __init__(self): super(BCELoss, self).__init__() self.bce nn.BCEWithLogitsLoss() def forward(self, preds, labels): if isinstance(preds, list): total_loss 0 for pred in preds: total_loss self.bce(pred, labels) return total_loss return self.bce(preds, labels)模型输出的多个侧输出list在和labels尺寸对齐时已经有interpolate操作所以可以直接计算损失。4.4 完整的图片背景去除推理代码接下来是核心的推理代码。以一个图像文件作为输入输出四张图原图、显著图、掩膜、合成到自定义背景的结果。import numpy as np import torch import cv2 from torchvision import transforms def load_model(model_path, devicecuda if torch.cuda.is_available() else cpu): model U2Net() state torch.load(model_path, map_locationdevice) if isinstance(state, dict) and state_dict in state: state state[state_dict] # 处理key前缀问题 new_state {} for k, v in state.items(): if k.startswith(module.): k k[7:] new_state[k] v model.load_state_dict(new_state) model.eval() return model.to(device) def preprocess_image(image_path, target_size320): img cv2.imread(image_path) img_rgb cv2.cvtColor(img, cv2.COLOR_BGR2RGB) h, w img_rgb.shape[:2] scale target_size / max(h, w) new_h, new_w int(h * scale 0.5), int(w * scale 0.5) resized cv2.resize(img_rgb, (new_w, new_h), interpolationcv2.INTER_LINEAR) # 用补边的方式保持长宽比并把尺寸固定为target_size delta_w target_size - new_w delta_h target_size - new_h top, bottom delta_h // 2, delta_h - (delta_h // 2) left, right delta_w // 2, delta_w - (delta_w // 2) padded cv2.copyMakeBorder(resized, top, bottom, left, right, cv2.BORDER_CONSTANT, value(0, 0, 0)) tensor transforms.ToTensor()(padded) tensor transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])(tensor) tensor tensor.unsqueeze(0) return tensor, img_rgb, h, w, (top, left, new_h, new_w) def predict_mask(model, tensor, device, use_d5True): with torch.no_grad(): outputs model(tensor.to(device)) if isinstance(outputs, list): if use_d5: # 融合所有侧输出 prob torch.mean(torch.stack(outputs[:6]), dim0) else: prob outputs[0] else: prob outputs prob torch.sigmoid(prob).cpu().numpy()[0, 0] return prob def remove_background(image_path, model, device, background_color(255, 255, 255)): tensor, original, h, w, pad_info preprocess_image(image_path) prob predict_mask(model, tensor, device) top, left, new_h, new_w pad_info # 去掉padding恢复到resize后的尺寸 prob_cropped prob[top:top new_h, left:left new_w] # 缩放回原图分辨率 mask cv2.resize(prob_cropped, (w, h), interpolationcv2.INTER_LINEAR) original_bgr cv2.cvtColor(original, cv2.COLOR_RGB2BGR) mask_3ch np.stack([mask] * 3, axis-1) # 白色背景去除 white_bg np.full_like(original_bgr, background_color, dtypenp.uint8) result (original_bgr * mask_3ch white_bg * (1 - mask_3ch)).astype(np.uint8) return original_bgr, mask, result这个推理代码有几个关键点输入图像不是直接resize到正方形那样会拉伸变形。我用了长边缩放加补边的方式既保持长宽比也让模型输入尺寸固定。推理得到的mask恢复到原图分辨率时用的是线性插值因为mask是连续概率值不是硬标签。合成背景时用的是alpha混合公式不是直接按阈值切一刀这样才能保留发丝级别的半透明过渡。5. 背景去除实战从DEMO到能用的抠图效果代码写完只是第一步跑出来效果好才算真的完成。我拿几张不同类型的图片做了实测这个环节很能说明问题。5.1 人像图效果惊艳但要注意边缘发丝第一张是典型的自拍人像背景是公园的树丛属于中等复杂度。模型推理出来的显著图质量相当高人的身体和头部区域几乎全白背景全黑边缘处有一圈灰色过度带。合成到白色背景后整体看起来比较自然。但要注意发丝区域如果原图中发丝和背景颜色接近或者背景存在和头发纹理相似的图案发丝部分会被误判进背景出现断发现象。这不是U2Net独有的问题是所有显著性检测模型的通病。5.2 商品图边缘锐度比预期好电商商品图通常是纯白背景主体是暖色产品。U2Net在这种图上的表现很稳定主体概率图非常实边缘质量高。这类图甚至不需要复杂的后处理直接阈值0.5就能得到不错的掩膜。如果商品是半透明的比如玻璃瓶、塑料袋U2Net会把整个半透明区域判定为前景无法区分透明物体内部的透过现象。透明和半透明物体的抠图需要专门的Matting模型U2Net做不到。5.3 多主体图显著图会倾向于把多个目标当做一个整体输入一张三个人并肩站着的照片U2Net输出的显著图会把三个人都标记为显著区域但三人之间的间距也被连带标记了因为网络学到的“显著性”是一个区域属性不是个体实例属性。这个特性决定了它的定位U2Net适合做“突出视觉主体”的粗分割不适合做“每个独立个体”的实例分割。如果业务需要多人分别抠图你需要在显著性分割后接实例分割模型。5.4 后处理技巧阀值选择与形态学操作拿到模型输出的概率图后一般会遇到两类情况概率图整体对比度很高前景接近1背景接近0中间过渡带很窄。这种情况直接阈值0.5就可以。概率图有些模糊边缘出现灰色羽化区。可以把它加进锐化处理或者对概率图做一次小范围高斯模糊再取阈值边缘过渡会平滑很多。还有一个经验如果你要做抠图合成不要用二值掩膜直接裁切那会让边缘像剪纸。而是把概率图本身当成alpha通道用。alpha通道的浮点灰阶过渡非常宝贵直接参与alpha混合能让结果自然得多。我在代码里已经用连续掩膜做了混合。但如果业务想要硬边缘PNG可以加一次形态学闭运算把边缘碎屑去掉kernel cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (3, 3)) mask_bin (mask 0.5).astype(np.uint8) mask_bin cv2.morphologyEx(mask_bin, cv2.MORPH_CLOSE, kernel, iterations2)5.5 视频里做背景替换性能与稳定性权衡U2Net做视频背景替换也是可行的但直接在每一帧上跑模型开销不小。我的实测经验用单一尺度推理在RTX 3060上处理1080p视频每帧大约0.1秒左右。如果要做实时至少需要把输入分辨率降到192或256并考虑用半精度推理。视频里连续帧的mask有时候会抖因为相邻帧的细微光照变化会影响分割结果。解决办法可以加一个时序稳定的后处理对mask做轻微的时间平滑。如果对实时性要求极高二阶段蒸馏一个小模型或者直接用U2Net的encoder部分做迁移学习是更实际的做法。6. 实测踩坑与效果提升U2Net没那么简单的几件事代码跑通不代表万事大吉。我在实战过程中踩过不少坑这里挑几个典型的记录下来。6.1 加载权重时遇到的key不匹配问题如果你从官方repo下载训练好的模型加载时经常遇到Missing key(s)或Unexpected key(s)的报错。原因有两种模型用了DataParallel训练权重key带module.前缀。模型定义和训练时不是完全一致。解决方式是写一个前缀清理函数我上面的代码里已经处理过了。遇到类似报错不要慌先打印几个key看看格式再做字符串替换即可。6.2 输入尺寸影响极大我在训练好的模型上测试时发现输入尺寸从320改成224精度下降明显尤其是细小物体和复杂边缘。这不是模型本身的问题而是在小分辨率下细节信息丢失了。反过来如果输入尺寸改成480边缘更好但显存占用大幅上升。实际工程中建议用插值方式将输入设置成一个固定batch_size可容忍的最大值配合自动混合精度AMP效果和速度可以兼得。6.3 光照变化对显著图的影响U2Net对光照是敏感的。同一件商品放在硬光下拍摄和散射光下拍摄输出显著图的边缘质量有明显差异。这说明训练数据本身的多样性决定了模型的鲁棒性上限。如果要做通用背景去除工具建议在推理前加一个简单的图像增强预处理归一化亮度、色彩校正能显著提升模型的稳定性。具体做法可以用自适应直方图均衡化CLAHE简化色彩增强6.4 数据增强怎么加才有效训练U2Net时随机翻转和随机裁剪是官方标配。我实验后还加了随机旋转不超过10度和色彩抖动对提升模型在自然图片上的泛化能力有帮助。但不要加得太狠旋转角度过大会破坏显著性分布的先验色彩抖动过强会让模型学到错误的颜色关联。6.5 生产环境部署要注意的事如果你打算把U2Net放到生产环境中有几点值得提前规划ONNX导出和TensorRT加速几乎是必然的。PyTorch直接用性能不够torch.onnx.export导出时要注意动态轴配置尤其是输入尺寸的动态变化。如果业务场景是固定分辨率输入比如手机端固定输出320x320可以把模型固定到某个尺寸导出时用静态shape性能更优。后端如果并发请求高不建议每个请求都加载一次模型。做一个常驻的推理服务利用模型预热、显存常驻、batch推理可以显著提升吞吐。6.6 U2Net和U2NetP的取舍U2Net还有一个轻量版本U2NetP参数量大约只有U2Net的十分之一左右精度下降不明显。如果你的目标是移动端或实时推理U2NetP的性价比很高。U2NetP的结构和U2Net几乎一样只是RSU各层的mid_ch通道数都缩小到了原来的1/4左右整个模型体积大幅降低。如果算力紧张直接换上U2NetP再蒸馏一下是个不错的折中方案。7. 更进一步从背景去除到更语义化的图像精细分割U2Net不是终点而是一个起点。显著性检测得到的显著图虽然能区分前景和背景但它不区分主体的更细粒度结构。比如一个人物你会把整个人作为前景却不知道四肢、头发、衣服具体在哪。如果你的业务需要的是“人像分部位”这种细粒度分割建议在U2Net的输出基础上再接一个精化网络。目前工业界比较成熟的组合方式是第一阶段用U2Net做主体区域定位。第二阶段在主体区域内做人像解析或语义分割。第三阶段用Matting算法细化毛发、边缘。这个多阶段的架构在很多商业直播、视频会议产品里已经在实际使用了。U2Net承担的是“先快速把目标区域框定”这个粗分割角色后续网络只需要在更小的区域内做精细分类计算量和样本难度都会大幅下降。这也正是我认为U2Net价值最大的地方——它是非常坚实的分割前端和特征提取器而不是一个只能端到端训练的“黑盒玩具”。我自己实际跑项目时最直观的感受是U2Net让“训练一个能用的分割模型”这件事的门槛降低了一个数量级。数据获取成本低、训练速度快、推理效果具备上线底气这三点就足以支撑它在众多分割算法中脱颖而出了。如果你正在做背景替换、商品抠图、内容裁剪或者图像合成相关的工作把U2Net放到你的工具箱里认真看看它输出的显著图长什么样你会在很多看似“需要专门分割模型”的任务上找到更轻量的解法。
返回列表