ARTICLE DETAIL

资讯详情

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

UNet多类别医学图像分割实战:从模型构建到后处理优化

UNet多类别医学图像分割实战:从模型构建到后处理优化 简介U-Net图像分割代码聚焦医学图像分割、语义分割与多类别分割任务适合需要设计或改进分割模型的学生、研究人员与算法工程师。网络采用对称收缩路径与扩展路径借助跳跃连接保留高分辨率特征对样本量有限的医学影像场景尤为适用。压缩包共31个文件大小仅16KB其中包括8个Python源码覆盖模型构建、数据集加载、训练、预测、混淆矩阵评估和数据增强等关键模块14个pyc编译文件便于快速加载与对比同时附带项目配置文件与说明文档。已有466人学习/下载。使用者可根据自身任务修改类别数、输入尺寸与数据读取逻辑将代码迁移至病灶分割、遥感语义分割等方向也能从训练脚本与评估流程中理解U-Net在语义分割任务中的工程化实现细节。1. 医学图像分割绕不开的起点UNet 与多类别分割的边界在哪里最开始接触医学图像分割时很多人的第一反应是跑通一个UNet。但真到了要处理多类别分割任务时才会发现问题不在网络结构本身而在数据标注方式、损失函数设计、评估指标选择这些环节。UNet之所以成为医学图像分割的基线模型核心在于它的U形对称结构——编码器逐级下采样提取语义特征解码器逐级上采样恢复空间分辨率再加上跳跃连接把浅层细节与深层语义融合这让它在中层特征不足的小样本医学数据上表现出色。本文要讲的不是单标签二分类分割而是多类别语义分割即每个像素被分到多个类别中的一个例如肝脏CT中区分背景、肝脏、肿瘤三类。与二分类不同的是多类别分割必须面对类别不平衡、小目标漏检、边界区域混淆三个典型问题。代码层面会涉及损失函数从BCE到CrossEntropy的转变、评估指标从Dice到mDice/mIoU的扩展、后处理时类别区域合并策略的改变。适合的人群是已经跑通二分类UNet、想转向多类别医学分割的工程师和研究人员以及在部署时被mDice卡在0.8以下需要找优化方向的人。这里给出的方案以PyTorch为框架覆盖从数据预处理到训练评估再到推理后处理的完整闭环。2. UNet 多类别分割的模型实现与关键参数对比2.1 编码器滤波器设计从64到512的扩展逻辑UNet原始结构以双卷积块为基本单元每个块包含两次卷积、BatchNorm和ReLU激活。编码器部分从64个滤波器开始每经过一次下采样翻倍直至512。这种滤波器数量递增的设计是为了在下采样过程中逐步提取更高层次的语义特征。多类别分割场景下类别数的增加意味着更细粒度的特征区分需求但盲目增加滤波器数量会带来显存压力。实验对比上看医学图像通常分辨率在512x512以内原始UNet的滤波器配置已经是精度与显存的较好平衡点。import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super(DoubleConv, self).__init__() self.conv nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue) ) def forward(self, x): return self.conv(x)上面是UNet中最基础的双卷积模块。padding1保证卷积不改变特征图尺寸。BatchNorm在医学图像小batch训练时非常关键因为医学数据集通常batch size只有4到8没有BN层的话深层网络很难收敛。编码器下采样使用nn.MaxPool2d(2)它不是可学习参数作用是把特征图尺寸减半同时保留显著特征并增大感受野。解码器上采样则使用转置卷积或双线性插值。转置卷积是可学习的能更好地恢复细节但容易产生棋盘效应双线性插值没有可学习参数输出更平滑。医学图像分割任务中考虑到器官边界往往需要平滑过渡双线性插值配合后续卷积修正更常见。跳跃连接是UNet的另一个核心设计。编码器第i层的特征图与解码器对应层的特征图在通道维度拼接这样浅层的边缘纹理信息可以直接传递到深层解码器弥补上采样过程中丢失的高频细节。多类别分割对边界定位要求更高跳跃连接的作用比二分类更明显。2.2 解码器与输出层设计类别数如何决定最终卷积解码器部分与编码器对称每次上采样后特征图尺寸翻倍、通道数减半再与跳跃连接传入的编码器特征拼接。最后一层用一个1x1卷积把通道数映射到类别数此时输出特征图的每个像素位置对应一个长度为类别数的向量这个向量经过softmax归一化后即为各类别的概率分布。class UNet(nn.Module): def __init__(self, in_channels, num_classes, features[64, 128, 256, 512]): super(UNet, self).__init__() self.ups nn.ModuleList() self.downs nn.ModuleList() self.pool nn.MaxPool2d(2) # 编码器 for feature in features: self.downs.append(DoubleConv(in_channels, feature)) in_channels feature # 解码器 for feature in reversed(features): self.ups.append(nn.ConvTranspose2d(feature*2, feature, kernel_size2, stride2)) self.ups.append(DoubleConv(feature*2, feature)) self.bottleneck DoubleConv(features[-1], features[-1]*2) self.final_conv nn.Conv2d(features[0], num_classes, kernel_size1) def forward(self, x): skip_connections [] for down in self.downs: x down(x) skip_connections.append(x) x self.pool(x) x self.bottleneck(x) skip_connections skip_connections[::-1] for idx in range(0, len(self.ups), 2): x self.ups[idx](x) skip_connection skip_connections[idx//2] if x.shape ! skip_connection.shape: x nn.functional.resize(x, sizeskip_connection.shape[2:]) concat_skip torch.cat((skip_connection, x), dim1) x self.ups[idx1](concat_skip) return self.final_conv(x)num_classes参数直接决定最后一层1x1卷积的输出通道数这个参数在实例化模型时传入例如肝肿瘤多类别分割传4背景、肝脏、肿瘤、血管胰腺分割传3。features列表控制每层滤波器数量默认是从64到512但实际使用中如果显存受限可以改为[32, 64, 128, 256]精度下降通常在2%以内。模型输出的是未经过softmax的logits。在PyTorch中计算交叉熵损失时nn.CrossEntropyLoss内部已经包含了softmax操作所以训练时直接把logits传给损失函数即可。但在推理阶段需要显式执行torch.softmax(output, dim1)才能得到概率分布再通过argmax得到最终的分割类别图。这里有个容易踩的坑跳跃连接拼接前要检查特征图尺寸是否一致。由于下采样过程中遇到奇数尺寸的特征图MaxPool后尺寸向下取整导致上采样后与原编码器特征图尺寸不一致。常见的做法是将编码器输入尺寸统一resize到偶数比如256x256或512x512或者在拼接前用插值对齐。3. 多类别医学图像分割的数据准备与类别权重计算3.1 标签格式与数据增强的类别一致性多类别医学图像分割的标签图与二分类不同。二分类标签是单通道二值图而多类别标签是单通道灰度图像素值对应类别索引0表示背景1表示第一类2表示第二类依次类推。因此数据增强必须保证原始图像和标签图的几何变换保持一致。随机翻转、旋转、缩放这些操作可以通过设置相同的随机种子实现但对弹性形变这类非线性变换最好使用同步变换的库函数。import albumentations as A from albumentations.pytorch import ToTensorV2 train_transform A.Compose([ A.RandomRotate90(p0.5), A.Flip(p0.5), A.ElasticTransform(alpha120, sigma120 * 0.05, alpha_affine120 * 0.03, p0.3), A.ShiftScaleRotate(shift_limit0.05, scale_limit0.1, rotate_limit15, p0.5), A.Normalize(mean(0.485, 0.456, 0.406), std(0.229, 0.224, 0.225)), ToTensorV2() ])在albumentations库中传入的标签图必须是[H, W]的numpy数组且和图像共享同一个随机变换状态这就是该库相比torchvision.transforms的优势。多类别分割中特别要留意ElasticTransform这类非线性形变因为它会导致像素级别的位移如果图像和标签不是同时变换类别边界会错位。训练自己的数据集时要注意标签是否连续。假设标注软件导出的标签是[0, 2, 5]中间缺少类别1和4此时需要在数据加载时做标签重映射将[0, 2, 5]映射为[0, 1, 2]。否则nn.CrossEntropyLoss会期望输出类别数等于最大标签值加1即6而模型输出只有3个通道导致维度不匹配。重映射的代码很简单np.unique(labels)获得实际存在的类别集合再构建旧值到新索引的映射字典。3.2 类别不平衡问题的权重计算策略医学图像分割中类别不平衡几乎是常态。例如肝脏CT中背景像素可能占95%以上肿瘤可能只占不到1%。如果直接使用默认的CrossEntropyLoss模型会倾向于把所有像素预测为背景因为这样已经能获得很高的准确率。解决方法是根据各类别像素占比计算权重让少数类的损失贡献更大。import numpy as np from collections import Counter def compute_class_weights(mask_paths, num_classes): pixel_counts np.zeros(num_classes) for mask_path in mask_paths: mask np.load(mask_path) # 假设mask以npy格式存储 for c in range(num_classes): pixel_counts[c] np.sum(mask c) total_pixels np.sum(pixel_counts) class_weights np.zeros(num_classes) for c in range(num_classes): class_weights[c] total_pixels / (num_classes * pixel_counts[c]) return torch.tensor(class_weights, dtypetorch.float32)权重计算采用中位数频率平衡法即每个类别的权重与像素占比成反比再除以类别数进行归一化。背景类的像素占比最大权重最小肿瘤类的像素占比最小权重最大。这个权重的取值范围通常从0.01到100不等为了让梯度更新更稳定可以对权重做进一步限制比如将最大值截断到10。除了中位数频率法还有一种常用的方法是直接取像素占比的负对数并归一化。两种方法在小目标分割上的表现差别不大但负对数法更平滑。如果训练过程中发现小目标类别偶尔被检测到、但mIoU波动剧烈可以尝试把两个方法的结果做加权平均。另一个角度是损失函数层面的解决方案。Focal Loss通过调制因子(1-p)^gamma来聚焦难分类样本其中gamma通常取2。在多类别分割中可以用CrossEntropyLoss的权重参数搭配Focal Loss组合使用。如何在代码中实现后续会涉及。3.3 数据加载器的类别维度处理多类别分割数据加载的关键在于标签的处理方式。PyTorch的DataLoader要求一个batch中的张量形状一致而医学图像的尺寸各不相同因此需要统一到固定尺寸。常见做法是训练时随机crop或resize到固定大小推理时采用滑动窗口。from torch.utils.data import Dataset class MedicalSegDataset(Dataset): def __init__(self, image_paths, mask_paths, transformNone): self.image_paths image_paths self.mask_paths mask_paths self.transform transform def __len__(self): return len(self.image_paths) def __getitem__(self, idx): image np.load(self.image_paths[idx]) mask np.load(self.mask_paths[idx]) # 统一尺寸 image cv2.resize(image, (256, 256)) mask cv2.resize(mask, (256, 256), interpolationcv2.INTER_NEAREST) # 标签重映射确保类别从0开始且连续 unique_labels np.unique(mask) remap {old: new for new, old in enumerate(unique_labels)} mask np.vectorize(remap.get)(mask) if self.transform: augmented self.transform(imageimage, maskmask) image augmented[image] mask augmented[mask] mask mask.long() # CrossEntropyLoss 需要 LongTensor return image, maskcv2.resize处理图像时用默认的双线性插值处理标签时必须指定INTER_NEAREST否则插值会引入不存在的新类别值。这是多类别分割代码中新手最容易犯的错误之一往往训练到一半发现损失变成NaN或类别数量不对排查到最后是标签resize方式错了。标签经过transform后是torch.Tensor最后一行转为long()类型。CrossEntropyLoss的target参数要求是LongTensor如果传入torch.float32会直接报错。数据加载器在多类别场景下还需要注意类别数的一致性——不同样本的标签类别数必须相同如果某张图的标注中缺了某一类np.unique返回的类别集会变少重映射后标签依然从0开始但整体类别数没变这一点和单个类别的存在与否无关因为模型的输出通道数是固定的。4. 损失函数选择、训练策略与 mDice 评估代码4.1 CrossEntropyLoss 与 Focal Loss 在多类别场景的取舍多类别分割的默认损失是nn.CrossEntropyLoss计算时对每个像素的类别预测概率取负对数再按权重加权平均。这种损失函数对各类别的梯度更新是显式的权重越大该类别每像素产生的梯度越大从而迫使模型更关注少数类。但CrossEntropyLoss存在一个现实问题它对像素级别的分类错误一视同仁不考虑区域结构因此在器官边界处容易出现锯齿状预测。Focal Loss在交叉熵基础上增加了调制因子对置信度高的样本降低损失权重让训练聚焦于难分类样本。原始论文中gamma取2效果最好。在多类别医学分割中Focal Loss对提升小目标类别如早期肿瘤的召回率有一定帮助代价是损失曲线收敛变慢且超参敏感。class FocalLoss(nn.Module): def __init__(self, alphaNone, gamma2.0, reductionmean): super(FocalLoss, self).__init__() self.gamma gamma self.reduction reduction self.alpha alpha def forward(self, inputs, targets): ce_loss nn.functional.cross_entropy(inputs, targets, weightself.alpha, reductionnone) pt torch.exp(-ce_loss) focal_loss (1 - pt) ** self.gamma * ce_loss if self.reduction mean: return focal_loss.mean() elif self.reduction sum: return focal_loss.sum() return focal_lossalpha参数直接复用上一节计算出的类别权重class_weights。如果数据集中各类别像素占比差异超过50倍建议使用CrossEntropyLoss加权重如果差异在10倍以内Focal Loss的调制效果更自然。实际工程中我先用带权重的CrossEntropyLoss训练20个epoch作为warm-up再切换为Focal Loss微调这样兼顾了收敛速度与边界精度。4.2 训练循环与学习率调度的细节决定收敛质量医学图像分割的训练batch size通常受限于显存特别是处理3D数据或高分辨率2D图像时。batch size为4时BatchNorm的统计量不稳定导致验证集Dice波动明显。一个常规做法是使用nn.SyncBatchNorm做跨卡同步或者改用InstanceNorm替代。训练循环中需要把梯度裁剪、学习率调度和验证评估紧密集成。梯度裁剪能避免损失出现NaN时梯度爆炸导致的训练崩溃。学习率调度使用ReduceLROnPlateau或余弦退火前者在验证指标停滞时降低学习率后者周期性重置学习率以跳出局部最优。from torch.optim.lr_scheduler import ReduceLROnPlateau model UNet(in_channels1, num_classes3).cuda() optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-5) scheduler ReduceLROnPlateau(optimizer, modemax, factor0.5, patience8) criterion nn.CrossEntropyLoss(weightclass_weights).cuda() scaler torch.cuda.amp.GradScaler() for epoch in range(epochs): model.train() for images, masks in train_loader: images, masks images.cuda(), masks.cuda() optimizer.zero_grad() with torch.cuda.amp.autocast(): outputs model(images) loss criterion(outputs, masks) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) scaler.step(optimizer) scaler.update() # 验证... val_mdice evaluate(model, val_loader) scheduler.step(val_mdice)torch.cuda.amp.autocast()混合精度训练能显著降低显存占用并加速训练。但需要注意当损失函数在fp16精度下出现下溢时如背景权重接近0导致梯度消失梯度裁剪和损失缩放会自动补偿。如果发现精度反而下降可以关闭amp或把GradScaler的init_scale调大。学习率设置方面医学图像分割从头训练时1e-4到3e-4比较合适使用预训练编码器时降到1e-5。AdamW相比Adam增加了权重衰减分离机制在UNet这种浅网络上防止过拟合的效果优于L2正则化。4.3 验证指标 mDice 与 mIoU 的代码实现与解释多类别分割最常用的评估指标是mDice和mIoU它们本质上是先计算每个类别的Dice或IoU再对所有非背景类别取平均。背景类别通常不参与计算因为背景像素占比过大把背景纳入平均会掩盖小目标的表现。def compute_mdice_miou(outputs, masks, num_classes, epsilon1e-6): outputs torch.softmax(outputs, dim1) preds torch.argmax(outputs, dim1) dice_scores [] iou_scores [] for c in range(1, num_classes): # 跳过背景 pred_c (preds c).float() mask_c (masks c).float() intersection (pred_c * mask_c).sum() union pred_c.sum() mask_c.sum() dice (2 * intersection epsilon) / (pred_c.sum() mask_c.sum() epsilon) iou (intersection epsilon) / (union - intersection epsilon) dice_scores.append(dice.item()) iou_scores.append(iou.item()) mdice torch.tensor(dice_scores).mean().item() miou torch.tensor(iou_scores).mean().item() return mdice, miouepsilon1e-6用于防止类别c在当前验证集的所有图像中都没出现时分子分母全为0的情况。多类别分割验证时还要注意类别的全局出现频率——假设某个类只出现在5%的训练样本中那么验证集的评估结果波动就会很大。常见的做法是在验证时同时记录每个类别的Volume Dice体积级别的Dice而不是像素级别也就是按整张3D图像计算Dice。这对医学影像任务尤其重要因为像素级Dice与临床观察的肿瘤体积偏差往往不一致。4.4 训练过程中的常见问题与排错方法训练UNet跑多类别分割时损失不下降或验证mDice卡在低值是高频问题。排查路径一般是固定的先看数据加载器输出的图像和标签是否正确再检查损失函数输入形状是否匹配然后观察预测结果的可视化。def visualize_predictions(model, val_loader, save_path, devicecuda): model.eval() images, masks next(iter(val_loader)) images, masks images.to(device), masks.to(device) with torch.no_grad(): outputs model(images) preds torch.argmax(outputs, dim1) fig, axes plt.subplots(3, 3, figsize(12, 12)) for i in range(3): axes[i, 0].imshow(images[i, 0].cpu().numpy(), cmapgray) axes[i, 0].set_title(Image) axes[i, 1].imshow(masks[i].cpu().numpy(), cmaptab10) axes[i, 1].set_title(Mask) axes[i, 2].imshow(preds[i].cpu().numpy(), cmaptab10) axes[i, 2].set_title(Pred) plt.savefig(save_path)可视化时如果标签图是torch.cuda.LongTensor直接matplotlib绘制需要用.cpu().numpy()转换。cmaptab10支持最多10个类别的离散着色类别数更多时使用nipy_spectral这类高区分度色图。多类别分割的可视化比二分类更复杂因为每种类别需要不同的颜色映射否则无法判断边界是否混淆。另一个常见问题是验证集mDice不错但测试集效果很差这通常与数据分布偏移相关。医学图像采集设备和参数不同导致同一器官在不同数据源中灰度值分布差异很大。常规做法是训练时加入强度增广random gamma correction、contrast adjustment、测试时使用test-time augmentation对输入做水平翻转和垂直翻转对输出概率取平均。后者虽然增加推理时间但在mDice上通常能提升1%到3%。5. 推理与后处理多类别分割从概率图到临床可用结果的关键步骤5.1 构建预测概率图的正确操作流程模型训练完成后推理阶段要做的工作比二分类复杂。单类别分割直接把概率图与阈值比较多类别则需要维护每个类别的独立概率图然后按像素位置取最大概率对应的类别索引。def predict_single_image(model, image, devicecuda, ttaTrue): model.eval() if isinstance(image, np.ndarray): image torch.from_numpy(image).unsqueeze(0).unsqueeze(0).float().to(device) with torch.no_grad(): if tta: # 水平翻转、垂直翻转的预测平均 outputs torch.softmax(model(image), dim1) outputs torch.softmax(model(torch.flip(image, dims[2])), dim1) outputs torch.softmax(model(torch.flip(image, dims[3])), dim1) outputs / 3 else: outputs torch.softmax(model(image), dim1) pred torch.argmax(outputs, dim1).squeeze(0).cpu().numpy() prob_map outputs.squeeze(0).cpu().numpy() # [C, H, W] return pred, prob_map翻转TTA的本质是对称性增强。医学图像中大多数器官的左右结构不完全对称比如肝脏在右侧垂直翻转对腹部CT的作用有限但不会产生负面影响。执行TTA时注意每次翻转后不需要翻转回原方向因为softmax输出是像素级概率分布翻转操作对每像素概率只做空间镜像不影响类别之间的对应关系。最终把三次概率图相加取平均得到更稳定的概率分布。多类别分割的一个潜在问题是类别间概率分布重叠。比如某个像素在背景和肝脏两类上的概率分别为0.51和0.48argmax会选择背景但这个像素很可能是边界区域。针对这种情况工程上可以引入类别先验约束例如肝脏和肿瘤具有空间包含关系即肿瘤一定位于肝脏内部。后处理时检查肿瘤预测区域是否与肝脏区域重叠如果完全不相交将其置为背景。5.2 多类别分割中的连通域分析与类别合并技巧得到argmax后的整数标签图后直接保存为.npy或.png在大多数场景下够用。但医学图像分割通常需要分析每个类别的独立区域——不同患者可能有多个肿瘤病灶需要统计每个病灶的体积、位置、最大径等量化参数这些指标与临床诊断直接相关。from scipy import ndimage def extract_connected_components(pred_mask, target_class, min_volume50): class_mask (pred_mask target_class).astype(np.int8) labeled_array, num_features ndimage.label(class_mask) components [] for i in range(1, num_features 1): component_mask (labeled_array i) volume component_mask.sum() if volume min_volume: coords np.argwhere(component_mask) bbox (coords[:, 0].min(), coords[:, 1].min(), coords[:, 0].max(), coords[:, 1].max()) components.append({ volume: volume, bbox: bbox, mask: component_mask }) return componentsmin_volume参数用来过滤噪声区域。多类别分割中小目标类别的输出往往在远离真实病灶的位置散布少量孤立的像素块这些通常来自模型的不确定性。根据类别不同最小体积阈值可以从50到500像素不等肿瘤类别建议更小的阈值。ndimage.label默认使用4连通还是8连通默认是4连通。医学图像中血管或细长结构在4连通下容易被切分为多个片段建议使用ndimage.label(input, structurenp.ones((3,3,3)))的方式显式指定6连通或26连通。在2D场景中用8连通3D场景中用26连通更符合解剖结构的连续性。5.3 类别合并与边界平滑的后处理策略医学图像分割的后处理还有一个隐藏需求——临床标注规范中通常只允许保留体积最大的连通域或者要求两个相邻类别不能出现锯齿状边界。最常见的后处理是条件随机场CRF平滑。训练好的UNet直接输出的概率图通常是空间平滑的但argmax后会出现细碎的误分类区域。CRF以像素强度为观察变量以类别标签为隐藏变量迭代优化能量函数能够有效消除这些噪声。import pydensecrf.densecrf as dcrf def apply_crf(image, prob_map, n_iter10, compat3): C, H, W prob_map.shape d dcrf.DenseCRF2D(W, H, C) # 将概率图转换为CRF需要的输入格式 unary -np.log(prob_map 1e-6).reshape(C, -1) d.setUnaryEnergy(unary) # 添加双边位置与颜色特征 img np.ascontiguousarray(image) d.addPairwiseBilateral(sxy20, srgb5, rgbimimg, compatcompat) # 迭代推理 Q d.inference(n_iter) pred np.argmax(Q, axis0).reshape(H, W) return predCRF的一元势函数来自UNet的负对数概率成对势函数基于像素位置与灰度值的相似性。sxy控制位置高斯核的带宽通常设为10到30值越大表示位置距离越远的两像素受抑制越强srgb控制外观一致性医学灰度图取值范围决定了该值在3到10之间。compat是成对势的权重一般取3到10。CRF会额外增加推理时间单张512x512图像约0.5秒到2秒在批处理任务中还算可以接受。后处理时的类别合并也需要编程处理。假设模型同时输出肝脏和肝肿瘤两个类别但标注格式是四类背景、肝脏、肝肿瘤、血管不需要合并。当模型有多尺度输出即同时预测粗粒度和细粒度的分割结果时需要用标签层次结构将粗粒度的某个类映射到多个细粒度子类或反向合并多个子类为一个父类。这类需求在带有器官和病灶双层标注的数据集中很常见合并时建议在概率层级操作而非预测标签层级操作——先合并子类的概率再取argmax比先取argmax再合并标签更平滑。5.4 代码工程化与推理加速的实用建议多类别分割模型的推理加速可以从模型结构、推理框架和批处理三个层面入手。模型结构层面把卷积替换为深度可分离卷积参数量可减少8倍左右mDice下降通常在1%以内。推理框架层面先将onnx模型通过polygraphy优化并做int8量化再用TensorRT推理。TensorRT对nn.BatchNorm2d的融合做得很好在A100上实测UNet推理速度能提升3倍。批处理层面多类别分割自然适合批量推理。与单类别不同多类别的批处理会产生更大的内存占用因为每个样本需要在通道维度上保存完整的概率图如果每张概率图是float32512x512x4的大小约为4MB批量32张就是128MB。推理时使用torch.cuda.amp.autocast()对UNet行混合精度推理可以同时控制显存和耗时。导出ONNX时有一个关于nn.Upsample的坑modebilinear在ONNX导出后在不同推理引擎中的行为可能有差异。建议导出时固定output_size而不是scale_factor这样可以减少引擎间的数值差异。另外多类别UNet的ONNX输出层是未softmax的logits部署到推理引擎时要在后续层自行加softmax操作有些推理框架并不支持直接对输出做argmax此时建议把softmax和argmax都放在CPU端完成GPU只负责计算logits。本文还有配套的精品资源点击获取
返回列表