简介:面向深度学习者与医学图像处理研究者的 Unet+Resnet 多尺度分割项目,以子宫颈细胞核分割任务为例,完成二分类语义分割全流程。已内置数据集、训练与预测脚本、训练好的权重,仅训练50轮即达全局像素准确率0.89、mIoU 0.72,加大训练轮数仍有提升空间,适合作为多类别分割的入门参考。
资源共804个文件,压缩包约113.33MB,包含387张jpg原始图像、383张png标注掩膜、8个Python源码、5个xml配置、3个txt说明、1个pth权重及readme文档;训练脚本采用随机缩放至设定尺寸0.5~1.5倍的多尺度策略,自动统计mask灰度并写入txt以设定Unet输出通道,便于扩展多分割任务。学习率使用cos衰减,run_results内可查看loss与iou曲线,训练日志还提供各类别的iou、recall、precision等指标。
已有228人学习下载,按readme操作即可完成训练与推理,适合快速跑通Unet分割项目或借鉴多尺度训练思路的读者。
1. 子宫颈细胞核分割的 2 分类实战:为什么 Unet+Resnet 和多尺度训练是标配组合
子宫颈细胞核分割是个典型的 2 分类语义分割问题:每个像素要么是核,要么是背景。可一旦换成真实的宫颈液基薄层细胞学切片,2 分类立刻变成折磨人的项目——细胞核大小可以从十几个像素跨到上百个像素,染色深浅、重叠程度、杂质噪声都会让普通 Unet 翻车。这个实战项目把 Unet 的编码器换成 Resnet,用残差结构把特征深度做上去,再靠预训练权重拉回医学数据量不足的劣势;多尺度训练则负责让模型在核尺寸变化面前保持稳定。它适合正处在“Unet 能跑通、效果总差一口气”阶段的从业者,也适合想往多类别分割扩展的团队。下面按做这类细胞病理分割的通用路径拆解:网络怎么改、多尺度训练怎么做、数据怎么喂、坑在哪。
2. 把 Unet 的编码器换成 Resnet:残差编码器怎么搭、预训练权重怎么接
2.1 为什么 Resnet 比继续加深 Unet 编码器更稳:短接、预训练与下采样路径
Unet 原始编码器是一串卷积 + 池化堆叠,结构简单,但深度一旦超过 5 层,梯度回传会明显吃力,模型容易停在局部最优。更麻烦的是医学数据量普遍不大,从头训一个深编码器,浅层特征收敛慢,分割边界会一直抖。
Resnet 在这里解决的不是“网络更深”这一个点,而是三个问题一起解决。
第一,残差短接让梯度有了一条从输出直通输入的通道,编码器堆到 34 层甚至 50 层都不容易梯度消失。第二,Resnet 在 ImageNet 上训好的权重可以直接加载,相当于给病理图分割模型一个“见过真实纹理”的初始化,这对小数据医学项目非常关键。第三,Resnet 的下采样路径是分段设计的:conv1 步长 2,后面 layer2/layer3/layer4 各做一次步长 2 的降采样,输出的特征图天然形成 1/4、1/8、1/16、1/32 的层级,正好能接上 Unet 解码器的逐级上采样。
换编码器不是把 Unet 编码器里的卷积替换成 BasicBlock 那么简单,还要处理一件事:Unet 原始跳跃连接通常保留 4 个尺度的特征,Resnet 同样有 4 个主干层级,但通道数不一样,解码器每一层的输入通道必须按 Resnet 实际输出调整。下面这个实现就是按这个思路写的。
2.2 用 PyTorch 搭 Resnet-Unet:完整网络结构与尺寸对齐
常见的做法是直接用 torchvision 里带预训练权重的 Resnet,把分类头丢掉,取 stage 输出喂给解码器。下面这段是一个能直接跑通的 Resnet34-Unet,输入输出都是单尺度,多尺度训练时输入尺寸可变,因为全卷积结构不受限制。
import torch import torch.nn as nn from torchvision import models class ConvBlock(nn.Module): """两次 3x3 卷积 + BN + ReLU,Unet 的基础卷积单元""" def __init__(self, in_ch, out_ch): super().__init__() self.block = nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) def forward(self, x): return self.block(x) class UpBlock(nn.Module): """转置卷积上采样 + 跳跃拼接 + 两次卷积""" def __init__(self, in_ch, skip_ch, out_ch): super().__init__() self.up = nn.ConvTranspose2d(in_ch, skip_ch, 2, stride=2) self.conv = ConvBlock(skip_ch * 2, out_ch) def forward(self, x, skip): x = self.up(x) x = torch.cat([x, skip], dim=1) return self.conv(x) class ResnetUnet(nn.Module): def __init__(self, num_classes=2, backbone='resnet34', pretrained=True): super().__init__() resnet = getattr(models, backbone)(pretrained=pretrained) # 编码器:保留 Resnet 的前五段输出 self.e0 = nn.Sequential( resnet.conv1, resnet.bn1, resnet.relu, resnet.maxpool ) # 1/4, 64 通道 self.e1 = resnet.layer1 # 1/4, 64 self.e2 = resnet.layer2 # 1/8, 128 self.e3 = resnet.layer3 # 1/16, 256 self.e4 = resnet.layer4 # 1/32, 512 # 解码器:通道数按 Resnet 实际输出调整 self.d4 = UpBlock(512, 256, 128) # 1/32 -> 1/16 self.d3 = UpBlock(128, 128, 64) # 1/16 -> 1/8 self.d2 = UpBlock(64, 64, 32) # 1/8 -> 1/4 # 回到原尺寸:1/4 -> 1/2 -> 原图 self.up1 = nn.ConvTranspose2d(32, 16, 2, stride=2) self.conv1 = ConvBlock(16, 16) self.up0 = nn.ConvTranspose2d(16, 8, 2, stride=2) self.conv0 = ConvBlock(8, 8) self.out = nn.Conv2d(8, num_classes, 1) def forward(self, x): s0 = self.e0(x) s1 = self.e1(s0) s2 = self.e2(s1) s3 = self.e3(s2) s4 = self.e4(s3) x = self.d4(s4, s3) x = self.d3(x, s2) x = self.d2(x, s1) x = self.up1(x) x = self.conv1(x) x = self.up0(x) x = self.conv0(x) return self.out(x)这段代码里最值得注意的点是尺寸对齐。e0 这一层包含了 Resnet 的 conv1 和 maxpool,所以输出是输入的 1/4;e1 之后还是 1/4,但通道被 layer1 展宽到了 64;e2 开始每过一个 layer 尺寸减半。解码器从 s4 起步,逐级拼回 s3、s2、s1,最后用两次转置卷积拉回原图。输入尺寸只要保证能被 32 整除就行,这也是后面多尺度训练里 base_size 要选 32 倍数的原因。
decoder 里没有拼 s0,是因为 s0 只经过 conv1+maxpool,语义质量不如 layer1 输出。如果显存够,把 s0 拼到 up1 之后的层能拉回一些边缘细节,但模型体积和显存占用都会上去,实际项目里一般先不拼。
2.3 灰度病理图和 RGB 预训练权重的通道处理
宫颈切片染色图读进来是 RGB,但很多病理扫描仪导出的是灰度 TIFF,或者你为了省内存把图转成了单通道。Resnet 预训练权重第一个卷积是 3 通道输入,单通道图直接 load_state_dict 会报 size mismatch。
我一般不会把单通道复制成三通道去硬凑,因为那样等于让前几层重复读同一份信息,预训练权重的通道结构没有充分利用。更常见的做法是把预训练 conv1 的权重在通道维求平均,改成 1 通道初始化:
from torchvision import models def load_pretrained_gray_conv1(model, backbone='resnet34'): resnet = models.__dict__[backbone](pretrained=True) pretrained_w = resnet.conv1.weight.data # [64, 3, 7, 7] gray_w = pretrained_w.mean(dim=1, keepdim=True) # [64, 1, 7, 7] # 替换模型第一个卷积并加载权重 model.e0[0] = nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3, bias=False) model.e0[0].weight.data = gray_w return model替换 conv1 后,原来 e0 里的 bn1、relu、maxpool 都还能继续用,不用改动其他层。要注意的是,torchvision 里 resnet.conv1 是独立模块,替换后模型其余预训练层依然可以正常加载。这一步做完,灰度图就能吃下 ImageNet 初始化,细胞核的边缘响应会比随机初始化明显更早收敛。
3. 多尺度训练实现:尺度抖动、乱序尺寸 batch 与三个调参点
3.1 细胞核分割里的多尺度到底解决了什么
一个宫颈细胞切片的扫描图里,细胞核直径随放大倍率和制片差异可以差 5 倍以上。小核可能只有 12 个像素,大核能到 80 个像素。如果训练时只用固定 512x512 的 patch 喂模型,感受野是固定的,模型学到的是“某个固定尺度下的核纹理”,换一批不同放大倍率的切片就翻车。
多尺度训练的核心不是把图放大几倍,而是让模型在每个 batch 里看到不同尺度的细胞核,逼着网络去学尺度无关的特征。常见做法有两种:一种是尺度抖动(scale jitter),即随机把图缩放到 0.75 到 1.25 倍再裁剪;另一种是 batch 内混入不同分辨率的 patch。后者对工程实现要求更高,需要一个能处理乱尺寸 batch 的 collate 函数。
我推荐的做法是两者结合:scale jitter 负责数据侧,collate padding 负责 batch 侧。这样既简单又能让每个 batch 天然存在尺度方差。
3.2 数据管线:随机缩放、随机裁剪和一个能对齐乱尺寸的 collate
下面这套 Dataset 结构是我在细胞核分割项目里一直用的模板,关键点都写在注释里了。
import random import cv2 import numpy as np import torch from torch.utils.data import Dataset from torch.nn import functional as F class ScaleJitterDataset(Dataset): def __init__(self, image_paths, mask_paths, base_size=512, scale_range=(0.75, 1.25)): self.image_paths = image_paths self.mask_paths = mask_paths self.base_size = base_size self.scale_range = scale_range def __getitem__(self, idx): img = cv2.imread(self.image_paths[idx]) mask = cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE) mask = (mask > 0).astype(np.uint8) # 统一成 0/1 h, w = img.shape[:2] # 1. 随机尺度:整张图先缩放 scale = random.uniform(*self.scale_range) new_h, new_w = int(h * scale), int(w * scale) img = cv2.resize(img, (new_w, new_h), interpolation=cv2.INTER_LINEAR) mask = cv2.resize(mask, (new_w, new_h), interpolation=cv2.INTER_NEAREST) # 2. 从缩放后的图里随机裁剪固定 patch if new_h > self.base_size: y = random.randint(0, new_h - self.base_size) else: y = 0 if new_w > self.base_size: x = random.randint(0, new_w - self.base_size) else: x = 0 img = img[y:y + self.base_size, x:x + self.base_size] mask = mask[y:y + self.base_size, x:x + self.base_size] # 3. 转 tensor,mask 保持 long 型 img = torch.from_numpy(img.transpose(2, 0, 1)).float() / 255.0 mask = torch.from_numpy(mask).long() return {'img': img, 'mask': mask} def __len__(self): return len(self.image_paths)scale_jitter 有两个细节不能省。第一,img 用 INTER_LINEAR,mask 必须用 INTER_NEAREST,否则 mask 会插值出 2、3 这种不存在的像素值。第二,缩放后如果比 base_size 小,直接保留原始尺寸,不要强行放大,因为小图放大没有新增信息,强行放大只会让模型学出模糊特征。
但上面这个 Dataset 返回的 img 尺寸不一致,而 PyTorch DataLoader 默认要求 batch 内所有 tensor shape 相同。解决方法是写一个 collate_pad:
def collate_pad(samples): imgs = [s['img'] for s in samples] masks = [s['mask'] for s in samples] max_h = max([i.shape[1] for i in imgs]) max_w = max([i.shape[2] for i in imgs]) img_batch = torch.zeros(len(imgs), 3, max_h, max_w) mask_batch = torch.zeros(len(imgs), max_h, max_w, dtype=torch.long) valid_batch = torch.zeros(len(imgs), 1, max_h, max_w) for idx, (img, mask) in enumerate(zip(imgs, masks)): h, w = img.shape[1], img.shape[2] img_batch[idx, :, :h, :w] = img mask_batch[idx, :h, :w] = mask valid_batch[idx, :, :h, :w] = 1.0 return { 'img': img_batch, 'mask': mask_batch, 'valid': valid_batch, }valid 这一路很多人会漏掉,但它是多尺度训练不翻车的关键。padding 区域不是真实图像内容,如果直接参与 loss 计算,模型会学到“把边缘 pad 区域预测成背景”,训练 loss 看似正常,验证时边界处会出现一条暗边。后面的损失函数会用到 valid 把 padding 区域排除掉。
3.3 多尺度训练的参数怎么调:scale_range、base_size 和类别像素
scale_range 我默认给 0.75 到 1.25,这是一个安全区间。如果切片本身分辨率高、核的尺寸跨度更大,就把范围放宽到 0.5 到 2.0,但要注意,放太大训练收敛会变慢,因为模型每次看到的同一张图形态差异过大。建议先用窄范围跑通,再逐步放宽。
base_size 必须选 32 的倍数,因为编码器最深层是 1/32。512 是显存和精度的平衡点;如果 batch size 撑不到 8,可以降到 384,但不要低于 256,否则大核的信息会被裁掉一半。
还有一个容易被忽略的参数是类别像素占比。细胞核在整张图里通常只占 5% 到 15%,如果随机裁剪经常裁到全背景 patch,模型会一直在学背景。我一般会加一个采样逻辑:每个 epoch 里强制 20% 的 patch 中心落在 mask 高亮区域附近,让细胞核样本不会被背景稀释。这个参数不在模型代码里,而是在 Dataset 的采样逻辑里,但它对收敛速度的影响比调学习率更明显。
4. 数据准备与 2 分类损失:从 mask 制作到 Dice/BCE 混合损失
4.1 标注转 mask 的格式和按片划分数据集的注意点
多类别分割项目里,标注一般来自病理医生在 WSI 查看器上画的多边形。导出时常见格式是 JSON,每个 ROI 是一串坐标点。转 mask 时要注意一个老坑:同一张图上多个 ROI 可能重叠,直接按顺序 fillPoly 会把后画的覆盖先画的。处理方式是把所有 ROI 按类别分组,先画背景再画核,或者对重叠区域做逻辑或操作。
import cv2 import numpy as np def polygons_to_mask(polygons, img_size, class_id=1): mask = np.zeros(img_size, dtype=np.uint8) for poly in polygons: pts = np.array(poly['points'], dtype=np.int32).reshape(-1, 2) cv2.fillPoly(mask, [pts], class_id) return mask这里的 class_id 对应 2 分类里的“核”。如果后面要做多类别扩展,比如把细胞质也标出来,class_id 改成 2、3 即可,网络输出层 num_classes 同步改。
数据集划分必须按“切片”或“患者”为单位,不能按 patch 随机划分。同一个切片里相邻 patch 的高度相似,如果训练集和验证集混着同一个切片的不同区域,验证 Dice 会虚高,等部署到新切片就现原形。做细胞病理项目,我通常按 WSI 文件粒度切分,保证一个切片的 patch 只出现在一个集合里。
4.2 用 Dice+BCE 混合损失解决细胞核占比过小
2 分类分割用 CrossEntropy 也能跑,但细胞核占全图比例太小,CE loss 会被背景类别主导,模型倾向于全预测背景,Dice 始终上不去。混合损失是这类任务的常规解:BCE 保证像素级梯度,Dice 直接优化区域重叠度,两者对类别不均衡都不敏感。
import torch import torch.nn.functional as F def dice_bce_loss(logits, targets, valid, alpha=0.5, smooth=1.0): # logits: [B,1,H,W] 原始输出,未过 sigmoid # targets: [B,H,W] 取值 0/1 # valid: [B,1,H,W] 1=真实区域, 0=padding bce = F.binary_cross_entropy_with_logits(logits, targets.float(), reduction='none') bce = (bce * valid).sum() / valid.sum().clamp(min=1.0) prob = torch.sigmoid(logits) inter = (prob * targets.float() * valid).sum() denom = (prob * valid).sum() + targets.float().sum() dice = 1 - (2 * inter + smooth) / (denom + smooth) return alpha * bce + (1 - alpha) * dicealpha 表示 BCE 的权重,0.5 是常见的默认值。smooth 加 1 是为了防止两个类别在极小区域内分母为 0,也顺便让 loss 曲线更平滑。注意 valid 必须乘在 BCE 和 Dice 的每个项里,尤其 Dice 的分子分母都要剔除 padding 区域,否则 batch 内尺寸不齐时 loss 会虚高。
4.3 训练超参和监控指标:一套直接照抄的参数表
下面这组参数是针对 512x512 输入、单卡 24G 显存能跑得动的配置。如果是 12G 显存,batch size 降到 4,学习率同步减半。
| 参数 | 推荐值 | 说明 |
|---|---|---|
| optimizer | AdamW | 比 Adam 收敛稳,配合 weight decay 更好 |
| learning rate | 1e-4 | 编码器预训练,建议比解码器低 3 倍 |
| weight decay | 1e-4 | 防止小数据过拟合 |
| batch size | 8 | 512x512 输入下的常见配置 |
| epochs | 80-120 | 细胞核数据集 1 万张以内,100 轮够 |
| scheduler | CosineAnnealing | 最后 20 轮学习率跌到 1e-5 以下 |
| 输入分辨率 | 512x512 | 必须能被 32 整除 |
| 混合损失系数 | alpha=0.5 | BCE 和 Dice 各占一半 |
训练时监控两个指标:训练集 Dice 和验证集 Dice。验证集 Dice 连续 15 轮不涨,先别急着加数据,去检查 mask 是否和图像对齐,这个坑比模型问题出现的概率大得多。
4.4 “多类别分割”和“2分类”不冲突:输出层扩展思路
这个项目的标题同时出现“多类别分割”和“2分类”,第一次接触的人容易绕晕。实际工程里并不矛盾:当前任务只把像素分成核和背景两类,所以是 2 分类;但代码里 num_classes 是独立参数,网络最后一个卷积输出通道是动态的。想从 2 分类扩到 3 分类(核、细胞质、背景),只需要把 mask 标注里增加一个类别值,然后把 num_classes 改成 3,损失函数从 sigmoid+BCE 换成 softmax+CrossEntropy 即可。
有一点要提醒:多类别扩展不是一个新模型,而是同一套 Resnet-Unet 骨架换输出头。细胞核和细胞质在边界处互相咬合,多类别训练时 Dice 会掉一点,但边界质量通常比单独训 2 分类更好,因为类别间的互斥信息被模型显式学到了。如果团队目标明确要做多个类别,建议第一天就把输出层设计成多通道,训练时从一个类别逐步加。
5. 细胞核分割训练避坑:5 条高频踩坑记录与排查方法
5.1 原图和 mask 尺寸对不上,训练第一步就崩
现象:DataLoader 刚跑第一个 batch,报 tensor size mismatch,或者 forward 到跳跃拼接维度对不上。
原因:mask 是从标注软件单独导出的,缩放时 mask 用了和原图不同的插值参数,或者标注软件自动加了边距。最常见的是读取 mask 时用了 cv2.imread(path, 1),导致 mask 变成三通道,尺寸虽然在,但 channel 维度直接冲突。
解决:进 Dataset 第一行就做 assert:
img = cv2.imread(img_path) mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) assert mask.shape[:2] == img.shape[:2], \ f"mask {mask.shape} != img {img.shape}"完成缩放和裁剪后,再断言一次,防止 resize 分支写错。
5.2 单通道输入加载 Resnet 预训练 conv1 报错
现象:加载预训练权重时报 state_dict size mismatch,卡在 conv1.weight 上。
原因:模型第一层是 nn.Conv2d(1, 64, 7),预训练权重是 [64, 3, 7, 7],通道数对不上。很多人以为把输入图复制三通道就完事,但那样模型第一层仍然是 3 通道,程序不报错,可如果你用单通道图去推理,输入张量维度又会不对。
解决:用前面 2.3 节的替换方法,把 conv1 改成 1 通道并取权重均值。注意替换后不要再整体 load_state_dict,而是把除了 conv1 之外的层单独加载,或者加载后再替换,顺序不能反。
5.3 loss 正常下降,Dice 一直是 0 的阈值问题
现象:BCE loss 从 0.8 降到 0.3,训练曲线很漂亮,但验证集 Dice 打印出来始终是 0。
原因:Dice 计算前没有对预测结果做 sigmoid 和阈值划分,直接把 logits 二值化成 0/1,所有负值都被置成 0,导致预测几乎全背景。还有一种情况是验证集 mask 值不是 0/1,而是 0/2 或 0/255,intersection 永远为 0。
解决:评估前统一做一次标准化,mask 读取后立刻执行mask = (mask > 0).astype(np.uint8)。计算 Dice 时,用pred = (sigmoid(logits) > 0.5)得到预测 mask,再和真实 mask 交并比。Dice 为 0 的问题一大半不在模型,在 mask 的数值范围。
5.4 多尺度训练后预测 mask 出现碎点和伪影
现象:训练完成,单尺度推理效果不错,但用多尺度预测时,mask 边缘出现很多孤立小点。
原因:多尺度 TTA 时,模型对同一张图在不同缩放下的输出做了平均,但小尺度下细胞核边界下采样次数多,边缘被平滑,平均后小尺度贡献的假阳性点残留了下来。另一个原因是 BN 在推理时用了训练集的统计量,尺度差异大时 BN 统计量漂移。
解决:多尺度评测时不要直接平均 logits,先对每个尺度的输出做一次小的形态学开运算,再平均。开运算可以清掉 2 像素以内的孤立点。如果仍频繁出现,检查推理时是否不小心把模型切到了 train 模式,导致 BN 用了 batch 统计量。
5.5 训练集验证集都高分,换一批新切片全崩
现象:交叉验证 Dice 0.87,部署到另一家医院或者另一台扫描仪采集的切片,Dice 直接掉到 0.5 以下。
原因:数据划分是按图而不是按切片,或者染色条件差异过大。病理切片染色是出了名的玄学,同一组织在不同实验室染色,色调差一个量级,模型学到的颜色特征在测试时失效。
解决:训练集划分强制按 WSI 文件粒度切分。同时在线做颜色增强:HSV 空间里对色调和饱和度做小幅扰动。更有效的办法是训练时随机转灰度,逼模型不要过度依赖染色颜色信号,把注意力拉回形态特征上。
6. 推理时把粘连核切开:watershed 后处理、TTA 与可视化验证
6.1 watershed 分离粘连细胞核
分割模型输出的概率图能分出核区域,但细胞核经常黏成一团,尤其在高密度区域,目标是两个核,预测出来是一个连通的 blob。这时距离变换 + watershed 是标准后处理方案。
import cv2 import numpy as np from scipy import ndimage as ndi def separate_nuclei(prob_map, threshold=0.5, min_distance=10): mask = (prob_map > threshold).astype(np.uint8) # 距离变换,核内部响应高,边缘低 dist = ndi.distance_transform_edt(mask) # 找局部极大值点作为 seed coords = ndi.maximum_filter(dist, size=min_distance) == dist markers, _ = ndi.label(coords) # watershed 分割,mask 限制在核内 labels = watershed(-dist, markers, mask=mask) return labelsmin_distance 是控制分裂粒度的关键参数。核直径平均 30 像素时,min_distance 给 10 到 12 比较合适;给太大,小核和背景一起被忽略;给太小,一个大核会被切碎成好几块。分离后统计每个 label 的面积,把小于 16 像素的碎片直接删掉,这类碎点通常不是核,是染色杂质。
6.2 TTA(测试时增强)和一页纸的可视化验证
模型部署前,我会用 4 倍 TTA 验证一次:原图、水平翻转、垂直翻转、旋转 90 度,四个结果都跑一遍,把 logits 对齐回原方向后平均。TTA 通常能带来 1 到 2 个点的 Dice 提升,但代价是推理耗时乘 4。如果项目对速度敏感,只保留水平翻转一种就够。
最后一步是可视化验证。每次训练完,把验证集里最好的和最差的几张图并排打印:原图、mask、预测 mask、叠加图。不要只看 Dice 数值,重点看边界处预测 mask 是否比标注更粗或更细。病理标注本身有主观性,两个医生标同一个核边界可能差 2 个像素,如果预测始终均匀偏粗,往往是数据集标注风格造成的系统偏差,不是模型问题。
我自己的习惯是每次启动训练前,先拿一个 batch 跑一次前向和反向,确认 loss 能降再挂全量训练。这个动作帮我省掉了大量“跑了一天发现数据加载是错的”的尴尬时刻。这套流程从 Resnet 编码器替换、多尺度训练到后处理分离,每一个环节都能单独验证,组合起来才是完整的细胞核分割落地路径。希望帮到你。
本文还有配套的精品资源,点击获取