☰
超声乳腺分割实战:基于BUSI数据集的UNet/ResUNet网页推理部署
2026/10/2 9:26:39 网站建设 项目流程

简介:这份资源是以超声乳腺疾病分割(BUSI数据集)为场景的医学图像分割项目,采用ResUNet与UNet双网络结构,支持在网页端完成可视化推理,适合正在学习医学影像语义分割流程的研究生、算法工程师及相关方向开发者。压缩包共900个文件,包括874张png图像、10个pyc、6个py、4张jpg、3个txt、1个readme、1个pth模型权重与1个json配置,总大小414.09MB,其中py脚本承载训练、验证与infer推理逻辑,png图像为数据集及可视化结果,pth和json则保存训练好的模型参数与评估记录。目前已吸引544人学习下载,代码可一键运行,训练部分内置ResUNet与UNet选择、余弦退火学习率调度和AdamW优化器,支持修改base-size适配大尺度训练;评估阶段输出dice、iou、recall、precision、f1、pixel accuracy等多项指标,并同步保存在runs目录下的json文件中。推理时执行infer脚本会在本地打开网页,上传图片即可看到分割效果,同时项目还保留了loss曲线、iou/dice曲线、学习率衰减曲线等可视化图表,便于直接分析模型训练过程与收敛情况。

1. 为什么超声乳腺分割要单独做一个网页版推理项目

基于网页版推理实现的ResUNet和UNet医学图像分割项目:超声乳腺疾病分割(BUSI数据集),拆开看其实只负责两件事:第一,用 BUSI 数据集把 UNet 和 ResUNet 训练成能分割乳腺超声病灶的模型;第二,把训练好的模型包装成网页服务,让不写代码的人打开浏览器上传一张超声图就能看到病灶掩码。医学图像分割的难点从来不只在网络结构,还在数据划分、预处理一致性和推理落地。命令行脚本只能自己用,网页版推理才算交付。下面按“数据准备、网络实现、训练验证、网页推理、自检验收”的顺序展开,适合刚开始跑 unet 模型、正在做医学图像分割项目的人。

2. BUSI数据集与超声乳腺分割任务:先搞懂要切的是什么

2.1 BUSI 数据集的真实结构

BUSI 数据集全称 Breast Ultrasound Images Dataset,是乳腺超声分割任务里很常见的基准数据集。文件通常按 normal、benign、malignant 三个类别目录存放,每张超声原图对应一张同名_mask.png掩码图。原图是超声设备输出的灰度图,掩码是黑白图,白色像素代表病灶区域,黑色像素代表背景。第一次拿到数据集不要直接训练,先做一次结构体检,否则后期各种问题都会归到“网络没调好”,其实是数据加载写错了。

一个典型的检查脚本长这样:

import os import cv2 data_root = "./BUSI/benign" images = [f for f in sorted(os.listdir(data_root)) if f.endswith(".png") and "mask" not in f] masks = [f for f in sorted(os.listdir(data_root)) if f.endswith("_mask.png")] # 检查原图和掩码是否一一对应 missing = [img for img in images if img.replace(".png", "_mask.png") not in masks] extra = [m for m in masks if m.replace("_mask.png", ".png") not in images] print("缺掩码:", missing) print("多掩码:", extra) # 统计掩码里有几个独立病灶 for m in masks[:5]: mask = cv2.imread(os.path.join(data_root, m), cv2.IMREAD_GRAYSCALE) _, binary = cv2.threshold(mask, 127, 255, cv2.THRESH_BINARY) num, _ = cv2.connectedComponents(binary) print(m, "独立病灶数:", num - 1)

这段代码有两个目的。第一,确认掩码文件名规则统一,避免后面 DataLoader 因为找不到对应文件直接崩掉。第二,用 connectedComponents 统计掩码里的独立连通域数量。BUSI 里经常出现一张掩码包含多个病灶的情况,后处理如果只保留最大连通域,就会漏掉另一个真实目标。先把这个数量统计出来,后续设计损失函数和后处理才有依据。

还有一个细节容易被忽略:掩码的像素值可能是 0 和 255,也可能因为标注工具导出变成 0 和 1。训练时统一阈值化成 0/1,不要直接在损失函数里用 255,否则 BCE 会被放大很多倍,模型很难收敛。另外,BUSI 的原图尺寸不统一,读取后先打印几组长宽,决定后续 resize 的基准尺寸。不要默认所有图都是 256x256。

2.2 超声图像为什么不能套用自然图像分割流程

自然图像分割里很多人习惯直接找 ImageNet 预训练权重,把 ResNet50 当作编码器。乳腺超声图像和自然图像差别很大,超声图灰度范围窄、对比度低、边界模糊,还充满散斑噪声。更麻烦的是病灶经常伴随声影,肿瘤内部灰度不均匀,边缘不闭合。从 ImageNet 迁移来的卷积核更擅长识别颜色和纹理,对超声图像并不友好。所以在 BUSI 这种医学图像分割项目里,网络设计往往从 UNet 这类医学分割结构开始,而不是盲目套语义分割大模型。

这并不意味着 UNet 落后。超声乳腺病灶和背景在灰度上差异不大,真正有用的反而是浅层边缘信息。UNet 的跳跃连接把下采样过程中丢失的高频细节直接送到解码器对应层,保留边缘定位能力。ResUNet 则是在这个结构上做 unet 模型改进,把每个普通卷积块替换成残差块,让信息能跨层传递。对小样本超声数据集,残差连接能缓解深层网络梯度消失,训练更稳。

还有一个容易踩的坑是输入通道数。BUSI 原图是灰度图,读出来是单通道,网络第一层输入通道要设成 1。很多代码从自然图像拷贝过来,默认输入 3 通道,结果训练和推理都不报错,但指标一直很差。处理超声图像时不要保留 RGB 三次重复,直接把灰度图变成[1, H, W]即可。如果设备导出的实际上本来就是三通道彩色超声图,也要显式转灰度,避免颜色通道引入噪声。

2.3 划分数据集时避免“同图不同帧泄漏”

小数据集上训练医学图像分割模型,最容易翻车的不是网络结构,而是数据划分。BUSI 里同一个病例可能有多张不同切面的超声图,这些图对应同一个病灶区域。如果随机按单张图片划分,同一个病例的图可能同时出现在训练集和验证集,验证指标会虚高。模型实际上记住了病例,而不是学会了分割。

正确做法是把文件名里的 case id 提取出来,按病例分组,再把整个 group 随机划分到 train/val/test。代码可以这样写:

import random from pathlib import Path def collect_pairs(root): pairs = [] for cls in ["normal", "benign", "malignant"]: cls_dir = Path(root) / cls for img in sorted(cls_dir.glob("*.png")): if "mask" in img.name: continue # 用文件名去掉后缀的方式配对掩码 mask = img.with_name(img.stem + "_mask.png") if mask.exists(): pairs.append({"image": img, "mask": mask, "case": img.stem.split("_")[0], "cls": cls}) return pairs pairs = collect_pairs("./BUSI") case_ids = list({p["case"] for p in pairs}) random.seed(42) random.shuffle(case_ids) train_cases = set(case_ids[:int(len(case_ids) * 0.8)]) val_cases = set(case_ids[int(len(case_ids) * 0.8):int(len(case_ids) * 0.9)]) test_cases = set(case_ids[int(len(case_ids) * 0.9):]) train_pairs = [p for p in pairs if p["case"] in train_cases] val_pairs = [p for p in pairs if p["case"] in val_cases] test_pairs = [p for p in pairs if p["case"] in test_cases] print("train/val/test:", len(train_pairs), len(val_pairs), len(test_pairs))

代码里img.stem.split("_")[0]是提取 case id 的常见做法。比如benign_1_2.png会得到benign_1。不同来源的 BUSI 命名可能略有差异,建议打印前 20 个 case id 确认。划分后最好按类别再检查一次分布。良性样本比恶性样本多,随机划分可能导致良性占满训练集,恶性样本在测试集表现很差。此时要对 case id 做分层 shuffle,而不是只对 pairs 做 shuffle。

训练集、验证集、测试集一旦划分完毕,整个实验周期不要再改动。固定random.seed(42)只是保证可复现,并不解决所有问题。同一个 case 的视频帧或不同切面,必须始终只出现在同一个集合里。验证时如果发现 Dice 高得离谱,最先怀疑的应该是数据泄漏,而不是模型能力强。

这一章看起来只是在处理数据,但医学图像分割项目里数据决定上限。超声乳腺分割的难点集中在边界不清和样本不均衡,UNet 和 ResUNet 都是对这种场景比较友好的结构,下一步就可以真正搭网络了。

3. UNet与ResUNet的选型与实现:把网络搭到能训练

3.1 UNet 结构:为什么跳跃连接对超声边界有用

UNet 由编码器、解码器和跳跃连接三部分组成。编码器通过“卷积 + 池化”逐步缩小特征图尺寸并增加通道数,解码器通过上采样恢复空间分辨率。如果没有跳跃连接,解码器只能依赖高层的抽象特征,浅层的边缘和细节会大量丢失。乳腺超声病变边界在灰度图上通常是几个像素宽的渐变带,这些高频细节保留在浅层特征图中,跳跃连接把它们直接拼到解码器同尺寸层,让网络同时参考语义信息和边界线索。

一个标准的卷积块定义如下:

import torch import torch.nn as nn class ConvBlock(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.block = nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1), nn.BatchNorm2d(out_channels), nn.ReLU(inplace=True), nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1), nn.BatchNorm2d(out_channels), nn.ReLU(inplace=True), ) def forward(self, x): return self.block(x)

这里两层 3x3 卷积都设置padding=1,特征图尺寸不变。BatchNorm 在 batch size 很小时统计不稳定,如果训练时 batch size 只有 4 甚至更小,建议把 BatchNorm 换成 GroupNorm。inplace=True能省一点显存,但网络结构生成时通常会保持原计算图,这个写法学起来不会有问题。

UNet 的编码器通道数一般从 64 开始,每次下采样通道翻倍,变成 64、128、256、512。底部再做一个 1024 通道的瓶颈层。显存有限时,不要硬上大 batch,先把第一层通道数改成 32。切换解码器时,输入通道数要记得加上跳跃连接传入的通道数,否则拼接后维度对不上。这些都是 unet 代码里最常见的报错点。

3.2 ResUNet 是把残差块塞进 UNet,而不是照搬 ResNet

ResUNet 的争议点在于到底是“改进的 UNet”还是“UNet 化的 ResNet”。常见做法是在 UNet 每个卷积块内部加入残差连接,整体结构仍然保持编码器、解码器、跳跃连接,所以本质上是带残差的 UNet。它不等于把 ResNet18 的下采样部分搬过来,因为 ResNet 没有对称的上采样解码器。

一个残差块可以写得很简洁:

class ResBlock(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.conv1 = nn.Conv2d(in_channels, out_channels, 3, padding=1) self.bn1 = nn.BatchNorm2d(out_channels) self.conv2 = nn.Conv2d(out_channels, out_channels, 3, padding=1) self.bn2 = nn.BatchNorm2d(out_channels) # 通道数不一致时用 1x1 卷积调整 self.shortcut = nn.Identity() if in_channels == out_channels else nn.Conv2d(in_channels, out_channels, 1) def forward(self, x): identity = self.shortcut(x) h = torch.relu(self.bn1(self.conv1(x))) h = self.bn2(self.conv2(h)) return torch.relu(h + identity)

参数说明:shortcut的作用是把输入张量对齐到输出通道数。输入输出通道相同时,直接使用恒等映射;不同时,用 1x1 卷积改变通道数,特征图尺寸不变。残差相加放在激活函数之前,两个卷积先提取特征,再和原始输入相加,最后一起经过 ReLU,这是残差网络里的标准写法。

把 ResBlock 拼成 UNet 时,每个 stage 内部可以连续放两个 ResBlock,下采样继续使用 MaxPool。这样改动最小,和标准 UNet 的编码器尺寸保持一致。如果一开始就加入注意力、空洞卷积等一堆改动,出了问题很难定位。先用标准 UNet 跑通训练,再把 ConvBlock 换成 ResBlock 做对比,是更稳妥的验证路径。

UNet 和 ResUNet 在 BUSI 上的取舍可以看下面这个表:

项目UNetResUNet
编码器 block普通双卷积带残差的双卷积
参数量少一些略多
梯度传导常规路径存在跳线短路
小样本收敛需要仔细调学习率通常稍微稳一点
显存占用较低略高

对 BUSI 这种规模的数据集,两种结构都能用。ResUNet 的改进意义体现在训练稳定性上,但不要指望 Dice 有质的飞跃。如果基础 UNet 已经能到 0.85,ResUNet 可能只提升零点几个点。选择的关键是硬件条件和实验目标。

3.3 训练参数和损失函数怎么定

BUSI 的分割目标是病灶区域,背景占比远大于前景,直接使用 BCE 会让模型倾向输出全背景,边界也很毛糙。一般做法是把 BCE 和 Dice Loss 组合在一起。Dice Loss 对前景和背景不平衡更稳健,BCE 则提供更平滑的梯度信号。两者相加可以互补。

一个常见的 dice_loss_from_logits 实现:

def dice_loss_from_logits(logits, target, smooth=1.0): prob = torch.sigmoid(logits) prob = prob.reshape(prob.size(0), -1) target = target.reshape(target.size(0), -1) # 计算每个样本的 Dice 再取平均 intersection = (prob * target).sum(dim=1) return 1.0 - (2.0 * intersection + smooth) / (prob.sum(dim=1) + target.sum(dim=1) + smooth)

说明:smooth是平滑项,防止分母为 0,也避免了训练前期数值震荡。target必须是 0/1 的浮点型,不能直接传入掩码文件读出的 255。手动实现 Dice Loss 时一定要展平后逐样本计算,不要在 batch 层面直接累加,否则不同样本的前景占比差异会互相干扰。

调用时组合方式通常如下:

bce = nn.functional.binary_cross_entropy_with_logits(logits, target) loss = bce + dice_loss_from_logits(logits, target)

如果发现小病灶漏检严重,可以把 Dice Loss 权重从 1.0 提高到 1.5。训练超参常见初始值是:输入尺寸 256x256,batch size 8,AdamW 学习率 1e-4,训练 100 个 epoch,配合 ReduceLROnPlateau。显存不够时优先缩小输入尺寸或减小 batch,而不是换一个更复杂的模块。unet 训练自己的数据集时,这个组合能覆盖大部分二分类分割需求。

4. 在BUSI上训练与验证:从数据生成器到 Dice/IoU 曲线

4.1 预处理和在线增强:让有限样本发挥更大作用

BUSI 的总样本量不大,直接训练很容易过拟合。常见做法是在线增强:随机水平翻转、垂直翻转、旋转、缩放。旋转时原图和掩码必须使用同一个变换矩阵,不能分开处理。OpenCV 的仿射变换可以同时对两张图操作,注意 ROI 外的填充值,原图填 0 问题不大,掩码填 0 表示背景,不要填 255。

预处理里最关键的是 resize 插值方式。原图下采样用INTER_AREA,可以避免高频混叠;掩码必须用INTER_NEAREST,否则原本 0/255 的二值掩码会变成 0~255 之间的灰色值,后续阈值化会产生多余边缘。归一化也要和推理端保持一致。比较简单的方案是把像素从 0~255 映射到 -1~1,因为超声图像灰度集中在中间范围,这样能让网络输入更平稳。

这里给出一个完整的 Dataset 类:

import numpy as np import torch import cv2 class BUSIDataset(torch.utils.data.Dataset): def __init__(self, pairs, img_size=256, train=False): self.pairs = pairs self.img_size = img_size self.train = train def __len__(self): return len(self.pairs) def __getitem__(self, idx): image = cv2.imread(str(self.pairs[idx]["image"]), cv2.IMREAD_GRAYSCALE) mask = cv2.imread(str(self.pairs[idx]["mask"]), cv2.IMREAD_GRAYSCALE) # 统一尺寸,掩码用最近邻保持二值 image = cv2.resize(image, (self.img_size, self.img_size), interpolation=cv2.INTER_AREA) mask = cv2.resize(mask, (self.img_size, self.img_size), interpolation=cv2.INTER_NEAREST) if self.train: if np.random.random() > 0.5: image = cv2.flip(image, 1) mask = cv2.flip(mask, 1) if np.random.random() > 0.5: image = cv2.flip(image, 0) mask = cv2.flip(mask, 0) image = (image / 255.0 - 0.5) / 0.5 mask = (mask > 127).astype(np.float32) # 返回 [1,H,W] 的张量 return torch.from_numpy(image).float().unsqueeze(0), torch.from_numpy(mask).float().unsqueeze(0)

代码说明:unsqueeze(0)把二维灰度图变成[1, H, W],对应模型的单通道输入。掩码二值化是mask > 127,不管版本是 0/255 还是 0/1,都能统一成 0/1。原图归一化后范围是大约 [-1, 1],如果训练时用 z-score,推理端也必须用同一组均值和标准差,不能改。

增强强度不要过大。超声图本身存在声影和器官变形,随机旋转超过 15 度可能让病灶相对位置失真。BUSI 训练里水平翻转和垂直翻转已经足够,再加小角度旋转即可。复杂的弹性形变在小数据集上未必稳定,建议先用简单增强跑一个基线,再决定是否升级。

4.2 训练自己的数据集:核心训练循环与参数

训练过程的关键不是把 epoch 跑完,而是保存验证集 Dice 最高的权重。很多人只保存最后一个 epoch,结果前面出现过更好的模型也丢掉了。下面是训练循环的核心片段:

def evaluate(model, loader, device): model.eval() dice_sum = 0.0 count = 0 with torch.no_grad(): for image, mask in loader: image, mask = image.to(device), mask.to(device) logits = model(image) # 注意:这里 logits 没有经过 sigmoid prob = torch.sigmoid(logits) pred = (prob > 0.5).float() intersection = (pred * mask).sum() denom = pred.sum() + mask.sum() dice = (2.0 * intersection + 1.0) / (denom + 1.0) dice_sum += dice.item() * image.size(0) count += image.size(0) return dice_sum / count model = ResUNet(in_channels=1, out_channels=1).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4) scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode="max", factor=0.5, patience=10) best_dice = 0.0 for epoch in range(100): model.train() for image, mask in train_loader: image, mask = image.to(device), mask.to(device) logits = model(image) loss = bce + dice_loss_from_logits(logits, mask) optimizer.zero_grad() loss.backward() optimizer.step() val_dice = evaluate(model, val_loader, device) scheduler.step(val_dice) if val_dice > best_dice: best_dice = val_dice torch.save(model.state_dict(), "best_resunet_busi.pth") print(f"epoch {epoch} best dice {val_dice:.4f} saved")

ReduceLROnPlateau的mode必须和监控指标对应,监视 Dice 时用max,监视 loss 时用min。很多新手在这里写反,导致学习率迟迟不降低。保存模型时只保存state_dict(),不要保存整个 model 对象。这样网页推理端可以自己定义模型结构然后加载权重,不受训练代码中模块路径影响。

DataLoader 可以设置num_workers=2或4加速读取。如果增强逻辑里有随机数,需要避免多个 worker 使用同一个随机种子导致数据重复。PyTorch 的 DataLoader 默认会对每个 worker 单独处理,不需要额外操心。验证集不要开 shuffle,数据顺序不会影响指标,但可复现性更好。

4.3 验证指标:Dice、IoU、PA 分别说明什么

二分类分割的验证指标主要看 Dice、IoU 和像素准确率 PA。像素准确率在超声分割里基本没意义,因为背景占比太大,模型全输出背景时 PA 也可能超过 90%。BUSI 上不同病例的病灶大小差别很大,只报一个总的 Dice 会掩盖小病灶漏检问题。建议同时打印 Dice、IoU,并按 normal、benign、malignant 三个类别分别统计。

一个直接可用的指标函数:

def calculate_metrics(pred, mask): pred = (pred > 0.5).astype(np.uint8) mask = (mask > 0.5).astype(np.uint8) inter = (pred & mask).sum() union = (pred | mask).sum() iou = inter / union if union > 0 else 1.0 dice = 2.0 * inter / (pred.sum() + mask.sum()) if (pred.sum() + mask.sum()) > 0 else 1.0 return dice, iou

实际使用中,Dice 会比 IoU 高一点,因为 Dice 对交集更敏感。不要只盯着验证集 Dice,还要随机挑几张测试集图片,把预测掩码叠加到原图上人工看边缘。有的模型指标不低,但边界多出几个像素,临床和演示场景里这种误差很影响观感。

如果预测结果出现大量散落的孤立噪声点,可以在后处理阶段用连通域过滤掉面积小于阈值的区域。这个操作会提升指标,但会掩盖模型本身的不足。写论文时如果要报告提升,必须把后处理写明;做网站演示时则无所谓,用户体验优先。另一个常见误用是把验证集指标当成测试集指标,如果前面按 case 划分没有做好,这个数字会虚高。先切好数据再训练,才能让指标有可信度。

5. 网页版推理实现与避坑:从 Flask 接口到 ONNX 加速

5.1 为什么网页推理要单独做,而不是直接让人跑 predict.py

模型训练好之后,如果只保留.pth文件,使用门槛很高。没有 GPU 的人要配置 PyTorch,还要读懂预测脚本里的参数,才能把一张超声图跑出掩码。网页版推理的目标是让操作者上传图片就能看到结果。这个环节容易失败,不是因为 Flask 难写,而是训练和推理的预处理不一致。训练时用了多卡、混合精度、在线增强,这些都不能直接搬进推理服务。

网页推理和训练推荐分开维护两个入口。训练代码重在可迭代,网页推理重在稳定。推理服务里不应该出现model.train()、反向传播、优化器,只保留前向计算。常见做法是先写一个独立infer.py,在命令行验证单张图片能输出掩码,再把这个流程封装成 Flask 接口。这样即使网页出问题,也能快速排除是模型还是接口的问题。

5.2 Flask 推理接口:模型加载、预处理、后处理

Flask 是单机网页推理里最常见的选项,代码量少,部署简单。BUSI 项目多半是内部演示或毕设展示,不需要高并发,Flask 足够。如果后续要做多用户同时上传,再换 FastAPI 或加队列也不迟。关键是把模型在模块加载阶段初始化,不要在每次请求时重新load_state_dict。

from flask import Flask, request, Response import torch import cv2 import numpy as np app = Flask(__name__) device = "cuda" if torch.cuda.is_available() else "cpu" # 模型类必须在这个文件里可导入 from model import ResUNet model = ResUNet(1, 1) state = torch.load("best_resunet_busi.pth", map_location=device) model.load_state_dict(state) model.to(device).eval() def preprocess(raw_bytes): arr = np.frombuffer(raw_bytes, np.uint8) img = cv2.imdecode(arr, cv2.IMREAD_GRAYSCALE) # 统一 resize 到训练尺寸,推理端不能改 img = cv2.resize(img, (256, 256), interpolation=cv2.INTER_AREA) img = (img / 255.0 - 0.5) / 0.5 x = torch.from_numpy(img.astype(np.float32)).unsqueeze(0).unsqueeze(0) return x @app.post("/predict") def predict(): file = request.files["image"] raw = file.read() x = preprocess(raw).to(device) with torch.no_grad(): logits = model(x) mask = torch.sigmoid(logits[0, 0]).cpu().numpy() # 二值化并转成 uint8,否则 PNG 编码会失败 mask = ((mask > 0.5).astype(np.uint8)) * 255 ok, encoded = cv2.imencode(".png", mask) if not ok: return "encode failed", 500 return Response(encoded.tobytes(), mimetype="image/png") if __name__ == "__main__": app.run(host="0.0.0.0", port=8000, threaded=False)

逻辑说明:torch.load(..., map_location=device)支持原来在 GPU 上训练的权重加载到 CPU 机器。preprocess返回的 shape 是[1, 1, 256, 256],固定 batch size 为 1。输出 mask 在二值化后乘 255,再用imencode编码成 PNG。前端只需要一个<img>标签,把请求发到/predict就能显示结果。

这个接口有三个容易忽略的地方。第一,resize 尺寸必须和训练一致,否则模型输入分布改变,掩码会变差。第二,上传图片可能是 BGR 或 RGBA,直接用IMREAD_GRAYSCALE可以规避通道问题。第三,threaded=False避免多线程下 PyTorch 推理互相抢占资源,如果并发量上来再单独处理。

提示:后端响应最好加上Cache-Control: no-store,否则浏览器可能缓存上一张图的掩码,换图后显示的还是旧结果。

5.3 模型导出到 ONNX,以及三个常见问题

如果部署机器不想安装完整 PyTorch,可以把模型导出成 ONNX,再用 onnxruntime 推理。导出时要先用model.eval()切到推理模式,固定一个输入尺寸导出,避免动态 shape 带来额外复杂度。训练时输入是 256x256,导出也保持这个尺寸。

导出代码可以这样写:

# 导出 ONNX 模型 model.eval() dummy_input = torch.randn(1, 1, 256, 256) torch.onnx.export( model, dummy_input, "resunet_busi.onnx", input_names=["input"], output_names=["output"], dynamic_axes={"input": {0: "batch"}, "output": {0: "batch"}}, opset_version=13 )

dynamic_axes只让 batch 维可变,宽高仍然固定。这样加载和推理都简单。ONNX Runtime 的推理接口如下:

import onnxruntime as ort import numpy as np session = ort.InferenceSession("resunet_busi.onnx", providers=["CPUExecutionProvider"]) def infer_onnx(x): x = x.numpy() output = session.run(None, {"input": x})[0] # 重新计算 sigmoid,把 logits 转成概率 mask = 1.0 / (1.0 + np.exp(-output)) return mask[0, 0]

常见问题一:torch.load报错提示找不到模块。原因是训练时保存了完整模型对象,类文件路径变了。解决方法是只用state_dict保存,网页端先实例化模型,再load_state_dict。常见问题二:导出 ONNX 后输出全是 0 或 1。原因是模型没有切到 eval 模式,BN 层在训练状态下使用了 mini-batch 统计。解决方法是导出前调用model.eval()。常见问题三:网页端处理大图很慢。原因是把原图直接 resize 到 512 甚至更大。BUSI 原图虽大,但分割任务对高分辨率不敏感,256 输入已经能拿到不错的效果。

网页端叠加显示可以交给后端做,把原图和 mask 用cv2.addWeighted合成在一起,再返回给前端。合成前要先记录原图尺寸,把 256x256 的 mask resize 回原图大小,否则叠加会错位。浏览器端不需要额外开 canvas,减轻了前端开发负担。真正做产品上线时,再考虑把预处理和后处理移到 GPU 上,但 BUSI 项目这个规模,CPU 推理完全够用。

6. 给网页推理做一次完整自检:单图、批量、异常图,一个都不能少

6.1 一个独立于训练代码的自检脚本

训练代码和网页推理代码是两套逻辑,很容易出现归一化参数不一致。我每次部署前都会写一个selfcheck.py,不通过 Flask,直接用同一套预处理函数读取测试集里的五张原图,把模型输出保存到outputs/目录。如果这批输出和训练时验证批次里看到的视觉效果一致,再启动 Flask。这个习惯能快速把模型加载、预处理、后处理分成三段排查。

自检脚本不需要太复杂,只要能稳定复现以下动作:加载权重、读取单张图、resize、归一化、前向计算、转 mask、保存可视化结果。我通常还会把原图和 mask 用np.hstack拼成一张图,一眼就能看出边缘是否对齐。这样即使后面改了训练尺寸,也能通过自检脚本发现网页端还留在旧尺寸。

6.2 给网页推理加一组“异常图”回归用例

有一次我被一张只有背景的超声图坑了。模型输出了一个全黑掩码,这本身没问题,但前端把黑图转成 PNG 后浏览器缓存了,第二次换图结果还是黑的。排查下来是 HTTP 缓存问题,解决方法是给响应加Cache-Control: no-store。如果提前准备一组回归用例,包括全黑输入、全白输入、纯噪声图、正常病灶图,这种问题能在几秒钟内暴露。

我把这些测试用例维护成一个固定列表,每次换权重或改输入尺寸都先跑一遍,跑完才敢把网页服务交给别人。网页推理项目真正的难点不在模型精度,而在输入分布变化和缓存这些细节。希望上面的实现和踩坑记录能帮你少走一段弯路,也希望你在部署时保留这个自检习惯,希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询