用GAN提升3D肝脏分割Dice:U-Net+对抗训练实战指南
2026/9/24 1:07:25 网站建设 项目流程

简介:面向医学图像分析与深度学习研究者的3D肝脏分割实现项目,基于生成对抗网络在Python环境中完成建模,并以Jupyter Notebook提供交互式实验流程。压缩包共12个文件,包含Python训练/预测/数据获取脚本、Jupyter Notebook示例、模型结构图、环境依赖与Shell运行脚本等,整体仅529KB,适合快速下载与本地复现。项目覆盖GAN生成器与判别器设计、3D卷积网络搭建、损失函数选择及Dice/Jaccard等分割指标评估,能够帮助读者深入理解生成对抗网络在三维医学图像分割任务中的实际应用。已有108人学习下载,适合具备一定深度学习基础、希望接触医学图像分割或GAN前沿研究的开发者。

1. 使用GAN做3D肝脏分割:为什么生成式方法能比纯分割网络多拿几个点的Dice

如果你手头正拿着这个“使用GAN进行3D肝脏分割_Python_Jupyter Notebook_下载.zip”,第一反应多半是解压、找 train.ipynb、改路径、点 Run。我劝你先别急着跑。这个包的核心不是某个现成权重,而是把“医学图像分割”改写成一个生成对抗问题:生成器从 CT 体数据里合成肝脏掩膜,判别器负责区分“真实标注”和“生成器输出”,两边较劲,最终把分割结果越逼越真。它的实际价值在于,当传统 U-Net 在低对比度边界和小病灶上反复翻车时,GAN 可以提供像素级损失给不了的局部对抗压力,让 Dice 再往上走一截。这篇笔记写给已经用 PyTorch 跑过至少一个分割项目、但没碰过 3D 医学图像 GAN 的从业者,也写给拿到 zip 不知道从哪个文件看起的人。下面按我跑通这类项目的顺序,把选型、代码、参数和坑一次讲完。

2. 任务定义与模型选型:3D肝脏分割到底难在哪、该选哪条GAN路线

2.1 肝脏分割的三道坎:低对比度、形态差异、标注误差

先解释“3D 肝脏分割”这个任务本身为什么难,不然你理解不了非得用 GAN 的理由。第一道坎是低对比度:肝脏在 CT 平扫中和周围组织(肌肉、肠道、胃壁)的密度差很小,单靠 HU 阈值几乎切不干净,边界往往要靠解剖位置的先验判断来补全。第二道坎是形态差异:不同患者的肝脏大小、形状、病灶位置、既往手术史都会让形状分布很宽,这导致分割模型特别容易过拟合到训练集的那几种形态上去。第三道坎是标注噪声:肝脏标注是医生在逐层切片上手动勾画的,层间不连续、边缘粗糙很常见,而这些标注噪声会被逐像素损失函数原样学进网络。

三道坎叠加在一起,就暴露出纯分割网络的短板。交叉熵或 Dice 这类损失函数对每个像素独立打分,然后把所有像素的误差平均成一个标量。这种平均化的监督信号“容忍”网络在每个位置都错一点点,却不关心局部形状是否合理——比如输出里多了一小截尾巴、边界上缺了个口子,平均损失可能只涨了千分之一。GAN 的判别器则不同:它见过大量真实标注的长相,能从结构上判断“这不像一个肝”。这正是生成对抗思路从图像修复、图像翻译一路延伸到医学分割的原因。

所以这个 zip 里真正值钱的不是某行“神奇代码”,而是那套把像素级损失和对抗损失揉在一起的训练目标。理解了这一点,后面调参时才不会把所有问题都怪到学习率头上。

2.2 为什么是GAN而不是再堆一层U-Net:对抗损失的边界收益

有人会问:纯分割网络效果不够,是不是把 U-Net 加宽加深就行?答案是不行。加深网络本质上还是让模型在同一类逐像素损失下做回归,它不知道“合理肝脏掩膜”这个群体的分布长什么样;而 GAN 用判别器学了一个高维度量,不只看每个像素对错,还看局部区域的形状、连续性、边界锐度。这才是结构性的损失。

注意,这不是说要把 Dice 损失扔掉。常见做法是“内容损失 + 对抗损失”叠加,生成器总损失写成两部分:

  • 内容损失:Dice Loss(或 Dice + 交叉熵),保证预测和标注在像素级尽量重合;
  • 对抗损失:判别器对“生成的分割图 + CT 输入”给出的真假分数,逼迫生成器输出整体结构更像真实标注。

GAN 的损失函数如果按参与者拆开,则是三份:生成器的对抗损失、生成器的内容损失、判别器上的二分类损失。这三份必须同时训练,哪一边太强都会把训练带崩——判别器太强,生成器得到的信息熵趋近于零,输出直接摆烂;生成器太强,判别器学不到有效梯度,对抗就形同虚设。这个平衡问题我放到后面避坑章重点讲。

实现上,这类任务多数采用条件 GAN(cGAN)的结构:把 CT 体数据当作条件输入生成器,生成器输出分割概率图;判别器的输入是“CT + 预测掩膜”或“CT + 真实掩膜”拼接后的双通道图像,输出一个反映真实性的打分。条件信息让判别器无法只看图像整体灰度分布就下结论,必须逐像素对照 CT 结构来判真假,这对医学图像来说很关键。

2.3 三种3D建模路线的对比:真3D、2.5D切片、伪3D,怎么选

拿到 zip 之后,第一个要确认的事是:包里生成器用的是 Conv3d,还是二维卷积 + 切片?这决定了你要准备多大的显存,也决定了代码能不能直接跑得动。我的经验是先把三种路线摊开对比,再决定要不要改造:

路线输入/输出形式显存占用优点缺点适用情况
真 3D整块 volume 或 3D patch,卷积核是 3D很高Z 轴上下文完整,层间连续性好训练慢,显存暴涨,要裁剪 patch显存 24G 以上,追求最终精度
2.5D 切片取相邻 N 层切片作为多通道输入层间信息有一定保留,可复用 2D 预训练权重本质还是 2D 卷积,层间连续性有限,输出层容易抖动显存 8-16G,性价比最高
纯 2D单层切片独立进出网络实现最简单,改动最小完全丢失 Z 轴信息,相邻切片的预测可能互相矛盾先跑通流程、验证 GAN 训练逻辑

我的建议很直接:如果这台机器只有一块 8G 或 12G 显存,不要一上来就抱着真 3D 不放。先把 zip 里的网络按 2.5D 的方式改造,把流程完整跑一遍,确认损失函数在降、Dice 在涨,再根据显存余量往真 3D 迁移。这比一步步在 OOM 报错里猜参数要靠谱得多。

3. 数据准备:从nii.gz到模型能吃到的训练样本

3.1 先把CT体数据读进来并做人肉质量检查

解压 zip 后,目录里一般会包含几个固定角色:数据目录(存放患者的 CT 体数据和标注)、训练用的 Notebook、模型定义文件、依赖清单。拿到手之后先别急着写 DataLoader,第一件事是打开几张样本,人眼确认数据和标注对不对得上。这是我最坚持的习惯,因为医学图像数据的方向矩阵、spacing、灰度范围各家医院差异极大,直接训练翻车率很高。

文件/目录常见内容先检查什么
data/volumes/患者 CT 体数据,常见 .nii.gz 或 .mhashape、spacing、方向矩阵是否一致
data/labels/与 volume 一一对应的肝脏掩膜取值是否为干净的 0/1,和 volume 是否同名
train.ipynb主训练流程数据路径是相对路径还是写死的绝对路径
requirements.txtPython 依赖PyTorch 版本、是否有 nibabel 等读取库

在 Jupyter Notebook 里第一个 cell 建议这样写,把样本和标签一起读出来核对:

import nibabel as nib import numpy as np # 读取 volume 和 label,第一件事是打印 shape 和标签取值 vol = nib.load("data/volumes/patient_001.nii.gz").get_fdata() lab = nib.load("data/labels/patient_001.nii.gz").get_fdata() print("volume:", vol.shape, vol.dtype) print("label:", lab.shape, np.unique(lab)) # 期望输出类似:(512, 512, 134) float64 / (512, 512, 134) float64 / [0. 1.] # 如果 label 取值范围是 [0, 255] 或 [-1, 1],先统一成 0/1 再做训练

这段代码的逻辑是:先确认两个文件的 shape 完全一致,再确认标签只有 0 和 1 两个取值。很多翻车现场都出在标记载体本身——例如某些标注工具导出的掩膜值是 255,或者把前景背景写反了,这类问题不提前发现,后面训练出的模型会带着系统性偏差。参数层面要注意 nibabel 读出来的 volume 默认是 float64,3D 体数据直接以这个 dtype 进网络会把内存撑爆,后面统一转 float32。

3.2 裁剪与归一化的代码:把肝区提出来再送进网络

一个典型的 CT 体数据是 (512, 512, 100~400) 的浮点数组,直接整块送进网络基本不现实,大多数实现会先把肝区裁出来。肝脏在腹部 CT 中只占整幅图像的一部分,四周是大量无关区域:空气、床板、肋骨、肌肉。把这些区域裁掉,既省显存,也能让归一化更稳。

我这里给一个预处理函数,固定在训练和推理两处复用,避免两边逻辑不一致:

def preprocess(vol, lab, margin=20, hu_min=-200, hu_max=200): # 用 label 的包围盒把肝区裁出来,四周留 margin 个像素 idx = np.where(lab > 0) z0, z1 = max(idx[0].min() - margin, 0), min(idx[0].max() + margin, lab.shape[0]) y0, y1 = max(idx[1].min() - margin, 0), min(idx[1].max() + margin, lab.shape[1]) x0, x1 = max(idx[2].min() - margin, 0), min(idx[2].max() + margin, lab.shape[2]) vol = vol[z0:z1, y0:y1, x0:x1] lab = lab[z0:z1, y0:y1, x0:x1] # HU 窗口裁剪:肝实质常用范围在 (-200, 200) 附近,具体按数据分布微调 vol = np.clip(vol, hu_min, hu_max) vol = (vol - hu_min) / (hu_max - hu_min) # 线性归一化到 [0, 1] return vol.astype(np.float32), lab.astype(np.float32)

参数说明有三处值得注意。margin 取 20,等于在肝区外围多留一圈组织当上下文,给卷积核提供边界判定的参考;hu_min 和 hu_max 控制窗宽窗位,-200 到 200 是很多公开肝脏分割项目在平扫 CT 上的常见取法,但如果你的数据是增强期 CT,碘剂让肝脏密度整体抬高,这两个值就要重新统计灰度直方图再定,别把它当金标准;最后转 float32 这一步是为了后续进入 PyTorch 时省一半内存。

3.3 用Dataset把数据和标签打包,注意增强一致性

预处理函数写完,就该把训练样本喂给 DataLoader 了。这里我给一个最小的 Dataset 骨架,核心是“随机裁剪 patch”和“volume 与 label 同步变换”这两点:

import torch from torch.utils.data import Dataset class LiverDataset(Dataset): def __init__(self, vol_paths, lab_paths, patch_size=(64, 256, 256)): self.vol_paths = vol_paths self.lab_paths = lab_paths self.patch_size = patch_size def __len__(self): # 每个患者多采样几个 patch,长度按训练集规模放大 return len(self.vol_paths) * 20 def __getitem__(self, item): idx = item % len(self.vol_paths) v = np.load(self.vol_paths[idx]) # 预先转成 npy,避免每次读 nii.gz l = np.load(self.lab_paths[idx]) # 随机裁剪:z、y、x 三个方向的起点,volume 和 label 必须用同一组下标 z = np.random.randint(0, v.shape[0] - self.patch_size[0]) y = np.random.randint(0, v.shape[1] - self.patch_size[1]) x = np.random.randint(0, v.shape[2] - self.patch_size[2]) v = v[z:z+self.patch_size[0], y:y+self.patch_size[1], x:x+self.patch_size[2]] l = l[z:z+self.patch_size[0], y:y+self.patch_size[1], x:x+self.patch_size[2]] # 转成 PyTorch 张量,加通道维度,float32 v = torch.from_numpy(v).unsqueeze(0).float() l = torch.from_numpy(l).unsqueeze(0).float() return v, l

这段代码里最容易出事的点就是裁剪坐标。z、y、x 三个随机数一旦生成,必须同时用于 volume 和 label 的切片,绝不允许分别调用一次随机函数。数据增强(旋转、翻转、弹性形变)同理,必须保证两个数组经历完全相同的变换;否则生成器会学到一个错位的映射,训练 loss 看着很低,画出图来边界却整体偏移。另一个小习惯是先把 nii.gz 预处理成 npy 缓存到磁盘,因为每次迭代都实时用 nibabel 读大体积文件,IO 会成为训练瓶颈。

4. 核心实现:在Jupyter Notebook里把GAN分割网络搭起来

4.1 生成器:U-Net风格编码解码结构

生成器沿用 U-Net 的编码-解码框架,是因为跳连结构能把下采样丢掉的边缘细节传回上层,这对分割任务特别关键。下面的代码是一个二维 U-Net 骨架,base 参数控制通道数,方便在显存和表达力之间权衡:

import torch import torch.nn as nn import torch.nn.functional as F class UNetGenerator(nn.Module): def __init__(self, in_ch=1, out_ch=1, base=32): super().__init__() self.e1 = nn.Sequential(nn.Conv2d(in_ch, base, 3, 1, 1), nn.ReLU()) self.e2 = nn.Sequential(nn.Conv2d(base, base * 2, 3, 2, 1), nn.ReLU()) self.e3 = nn.Sequential(nn.Conv2d(base * 2, base * 4, 3, 2, 1), nn.ReLU()) self.d2 = nn.Sequential(nn.ConvTranspose2d(base * 4, base * 2, 4, 2, 1), nn.ReLU()) self.d1 = nn.Sequential(nn.ConvTranspose2d(base * 2 + base * 2, base, 4, 2, 1), nn.ReLU()) self.out = nn.Conv2d(base + base, out_ch, 1) def forward(self, x): e1 = self.e1(x) e2 = self.e2(e1) e3 = self.e3(e2) d2 = self.d2(e3) # 跳连前检查尺寸,下采样过程可能让特征图差 1 个像素,用插值对齐 if d2.shape[2:] != e2.shape[2:]: d2 = F.interpolate(d2, size=e2.shape[2:], mode="bilinear", align_corners=True) d1 = self.d1(torch.cat([d2, e2], dim=1)) if d1.shape[2:] != e1.shape[2:]: d1 = F.interpolate(d1, size=e1.shape[2:], mode="bilinear", align_corners=True) return torch.sigmoid(self.out(torch.cat([d1, e1], dim=1)))

这个结构做了两次下采样,输入尺寸如果接近 256×256,显存占用非常温和。forward 里两个 if 判断是我后来补上的,因为卷积和转置卷积在 stride=2 时对奇数尺寸不友好,特征图尺寸会差一个像素;直接暴力拼接会让 torch.cat 报维度错误。输出层接 sigmoid,把分数压到 (0,1) 区间,和标签的 0/1 对齐。

如果你确认包里的模型是 3D 版本,最简单的改法是把所有 Conv2d 换成 Conv3d、ConvTranspose2d 换成 ConvTranspose3d,输入输出维度加一维。同时 patch 的 Z 轴深度要相应调小,显存一般会乘上 Z 方向切片数的倍数,先试 16 层。

4.2 判别器:PatchGAN对局部真假判定

判别器我建议用 PatchGAN,而不是输出单一标量的普通二分类。原因是单一标量只告诉生成器“整张图看着像不像”,生成器可以靠整体糊弄过关,局部边界仍然不行;PatchGAN 把特征图划分成若干小块,每一块都输出一个真假分数,强迫局部结构都逼真。肝脏分割最看重的边界细节,恰好是这种局部约束收益最大的地方。

class PatchDiscriminator(nn.Module): def __init__(self, in_ch=2, base=32): super().__init__() self.conv = nn.Sequential( nn.Conv2d(in_ch, base, 4, 2, 1), nn.LeakyReLU(0.2), nn.Conv2d(base, base * 2, 4, 2, 1), nn.BatchNorm2d(base * 2), nn.LeakyReLU(0.2), nn.Conv2d(base * 2, base * 4, 4, 2, 1), nn.BatchNorm2d(base * 4), nn.LeakyReLU(0.2), nn.Conv2d(base * 4, 1, 4, 1, 1), ) def forward(self, x): # x: (B, 2, H, W),两通道分别是 CT 切片和分割掩膜 return self.conv(x)

输入通道 in_ch=2,第一个通道放 CT 切片,第二个通道放真实标注或生成器输出。注意判别器内部不要用 ReLU,改用 LeakyReLU(0.2),因为 ReLU 会把负区间全部截断,判别器容易直接“死掉”——所有输出都是正数,真假无法区分。BatchNorm2d 放在中间层能稳定训练,但如果发现判别器 loss 下降过快,把它删掉有时反而更平衡。

4.3 损失函数:对抗损失与Dice损失怎么配比,加不加feature matching

损失函数是 GAN 分割项目里最微妙的部分。常见做法是内容损失用 Dice Loss,对抗损失用二分类交叉熵,两者按权重相加。Dice 损失管像素级重合,对抗损失管结构真实度。对于一个分割项目,Dice 权重通常是 1.0,对抗权重从 0.1 起调;太小了 GAN 没有存在感,太大了生成器会被对抗信号带偏,只学骗判别器而忽略边界精度。

feature matching 这个技巧(在 Salimans 等人的 Improved Techniques for Training GANs 里提出)可选,它的思路是:让生成器的中间层特征逼近真实样本经过同一个判别器时的中间层特征,相当于给生成器修了一条通往真实分布的缓坡。我在训练不稳定的时候会加上,代码上就是取判别器某一层输出,对真实和生成两组特征求 L1 距离。

def dice_loss(pred, target, smooth=1e-5): # pred 是 sigmoid 之后的连续值,target 是 0/1 硬标签 pred = pred.reshape(pred.size(0), -1) target = target.reshape(target.size(0), -1) inter = (pred * target).sum(dim=1) return 1 - (2 * inter + smooth) / (pred.sum(dim=1) + target.sum(dim=1) + smooth) # 生成器总损失 = 内容损失 + 对抗损失 # g_loss = lambda_dice * dice_loss(fake, lab) + lambda_adv * bce(d_fake, ones) # 判别器总损失 = 对真实标签判错 + 对生成结果判对 # d_loss = (bce(d_real, ones) + bce(d_fake, zeros)) / 2

smooth 参数要加,防止分母为零导致 loss 变成无穷大。Dice 的输入是连续概率图,不是先对输入做阈值二值化再算,因为取阈值之后梯度断掉了,无法回传。如果你更习惯交叉熵,也可以把 Dice 换成 BCE——但分割任务里前景背景像素严重不均衡,纯 BCE 容易被背景主导,所以我默认 Dice。feature matching 那一项别分太大权重,0.1 以下尝到甜头就够了,加多了会压制生成器的多样性。

4.4 训练循环的骨架:两个优化器、三段forward、怎么记日志

训练循环是整套代码里最不允许写错的部分。两个优化器,两段 backward,顺序必须严格:先更新判别器,再更新生成器。更新判别器时,生成器的输出要 detach 掉,否则梯度会顺着计算图流回生成器,等于同时改了两边,训练直接失控。

for epoch in range(epochs): for vol, lab in dataloader: vol, lab = vol.cuda(), lab.cuda() ones = torch.ones(vol.size(0), 1).cuda() zeros = torch.zeros(vol.size(0), 1).cuda() # 一、先训练判别器 fake = gen(vol) d_real = disc(torch.cat([vol, lab], dim=1)) d_fake = disc(torch.cat([vol, fake.detach()], dim=1)) d_loss = (bce(d_real, ones) + bce(d_fake, zeros)) * 0.5 opt_d.zero_grad() d_loss.backward() opt_d.step() # 二、再训练生成器 fake = gen(vol) d_fake = disc(torch.cat([vol, fake], dim=1)) adv_loss = bce(d_fake, ones) # 希望生成结果被判为真 content_loss = dice_loss(fake, lab) # 希望像素级重合 g_loss = lambda_adv * adv_loss + lambda_dice * content_loss opt_g.zero_grad() g_loss.backward() opt_g.step()

d_fake 在第二步没有 detach,这是故意的——生成器的梯度就是要从判别器的输出流回生成器。两个优化器分别只有一组参数,所以不会互相覆盖。lambda_adv 我一般先设 0.1,lambda_dice 设 1.0,然后每跑完一个 epoch 盯着验证集 Dice 和生成器输出的可视化图,而不是只看 loss 数字。GAN 的 loss 曲线并不能直接反映分割质量,判别器和生成器互相玩猫鼠游戏,loss 可能一路乱跳但 Dice 在涨,也可能两边 loss 都很稳但预测一团糟。所以每 N 个 step 保存一对 (CT, label, prediction) 的切片叠加图,比任何指标都直观。

5. 训练常见问题排查:翻车现象、根因与解决方案

5.1 判别器loss秒变0,生成器输出一坨灰

现象:训练不到 50 步,d_loss 降到 0.000x,TensorBoard 里生成的预测图变成一片均匀的灰色。原因:判别器太强。常见于判别器网络比生成器深、或者没加 BatchNorm、或者真实标注与生成输出的差异过于明显,导致判别器一上来就有碾压级的分类能力,梯度对生成器来说几乎没有信息量。解决:先改结构,给判别器中间层加 BatchNorm,把通道数减半;再改训练节奏,让生成器每步更新两次、判别器更新一次;也可以给真实标签乘 0.9 做标签平滑,降低判别器的置信度。我习惯先把 lambda_adv 压到 0.1,等稳定了再逐步抬起来。

5.2 显存OOM:3D卷积直接把显卡爆掉

现象:patch 切到 (64, 256, 256) 后,训练进行到第二个 epoch 附近报 CUDA out of memory。原因:Conv3d 的中间特征图张量体积远大于同尺寸 2D 卷积,加上 U-Net 跳跃连接会在前向过程中缓存大量中间结果,显存直接翻倍。解决:三个手段按性价比排——把 patch 的 Z 轴从 64 降到 32 或 16;把真 3D 卷积改成 2.5D,用相邻 5 层切片拼成通道;再打开混合精度训练,也就是 torch.cuda.amp 把前向和反向降到 float16,显存能省近一半。第一次跑通流程不建议硬上真 3D。

5.3 训练结束Dice反而比纯U-Net低

现象:loss 一直在降,但验证集 Dice 只有 0.75,同数据下纯 U-Net 能到 0.84。原因:这是 GAN 分割最容易翻车的地方——对抗损失把生成器带偏了,它学到的是“骗过判别器”的纹理,而不是精确的肝脏边界。像素级内容损失被对抗信号稀释,边界精度实际在退步。解决:把损失策略改回“内容为主、对抗为辅”,lambda_dice 提到 5.0,lambda_adv 压到 0.05;或者更干脆,前 80% 的 epoch 只跑内容损失,让生成器先拿到一个像样的初解,后 20% 再打开对抗损失做精修。这种“后悔药”式的分阶段训练在不少公开复现里都有效,值得先试。

5.4 volume和label错位,生成器学到错误映射

现象:训练 loss 很低,但把预测叠加到 CT 原图上,发现边界整体偏移一两毫米,且偏移方向固定。原因:数据增强或裁剪时,volume 和 label 用了两套独立的随机变换。比如对 volume 做了 90 度旋转,label 却没转,生成器学到的是错位映射。解决:增强逻辑必须写成一个函数,输入是 (volume, label),输出也是 (volume_t, label_t),所有旋转、翻转、缩放操作在函数内部用同一组参数实施,绝不允许分成两行代码独立处理。这个错误不会暴露在 loss 里,但会在验证集可视化时原形毕露。

5.5 Jupyter里Dataloader卡死,训练不推进

现象:cell 一直在转圈,GPU 利用率 0%,但进程没报错。原因:两个常见点。第一,Windows 下 DataLoader 的 num_workers 设成大于 0,多进程 fork 可能和 Jupyter 内核冲突;第二,Dataset 的getitem里有死循环,比如随机裁剪时 patch_size 大于输入尺寸,np.random.randint 的区间变成负数,代码在 while 里空转。解决:先把 num_workers 设成 0,跑通后再尝试调成 2;在getitem里对输入尺寸做 assert,保证大于 patch_size;更彻底的做法是预处理一步到位,直接生成 npy 缓存文件,每个 epoch 不再重复做裁剪和归一化。

6. 把Notebook变成能复现的结果:验证指标、模型导出与一次推理工作流

训练结束不是终点,能复现才是。我见过太多人把代码留在 Jupyter 里,一周后回来连自己都看不懂哪个 cell 对应哪个实验。我的习惯是训练一结束就立刻做三件事:算硬指标、存 checkpoint、抽一个推理函数。

验证指标上,Dice 是肝脏分割最通用的指标,但它对边界误差不敏感,所以加上 Hausdorff 距离(HD95)看边界偏差。Dice 高但 HD95 大,说明总体重合度还行、边界却有一处明显鼓包,这正是 GAN 分割常见的失败模式,只看 Dice 根本发现不了。一次推理的标准流程包含三行固定逻辑:读 nii.gz → 预处理(裁剪、归一化,和训练时完全一致)→ 转张量进模型 → 把输出 resize 回原始坐标空间。注意推理时不需要随机裁剪,而是裁剪到固定包围盒,预测完再映射回去。模型导出用 torch.save(model.state_dict(), path) 保存生成器就够,不需要存整个优化器状态;如果你还想继续微调,才需要把 optimizer 的 state_dict 一起存。

def predict(vol_path, model, device): # 推理:和训练预处理保持完全一致,再用包围盒还原坐标 vol = nib.load(vol_path).get_fdata().astype(np.float32) vol, _ = preprocess(vol, np.zeros_like(vol)) # 推理时没有 label,只做归一化 patch = torch.from_numpy(vol).unsqueeze(0).unsqueeze(0).float().to(device) with torch.no_grad(): pred = model(patch) return pred.squeeze().cpu().numpy() > 0.5

每次都手动验证的话,几行功能代码散落在各个 cell 里,很容易出现训练和推理预处理不一致。我现在的做法是把 preprocess 和 predict 抽到一个 model.py 文件里,Notebook 只负责调用,这样训练、验证、部署共用一套逻辑,不会再出现“训练时归一化到 0-1,推理时忘记归一化”这类低级事故。

另外固定随机种子也很重要。PyTorch 的卷积初始化、DataLoader 的 shuffle、cudnn 的自动调优都会引入随机性,不固定种子,同一份代码跑两次结果可能差一截。开头的 cell 加上 torch.manual_seed(0)、np.random.seed(0)、random.seed(0),并把这三个种子写进实验记录文件,这样才能保证哪天回头复现实验时,不会对着一个不可复现的数字发愁。我以前图省事总是直接双击 Run All,后来发现哪个结果都复现不出来,才把数据路径、随机种子和损失版本号都记进实验表里,这个习惯省下的时间远多于它花掉的时间。希望帮到你。

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

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

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

立即咨询