如果你在做医学图像分割,特别是病灶这一类小目标,应该会有同感:真实标注样本少,模型很容易在训练集上过拟合,到了新的病例上病灶形状、大小、位置一变,分割效果就往下掉。OTLesMix 这个方案,核心思路是用 Wasserstein 重心和最优传输映射来做合成病灶生成。它不只是在像素层面把两张图混一混,而是先估计真实病灶的形状-位置分布,再通过传输映射生成新的病灶样本,让增强后的训练集在形状和位置上更接近真实场景。这篇文章会从数学直觉讲起,然后给出一个可以在本地跑通的流程,包括环境准备、最小输入格式、关键参数、质量判断和常见坑点。适合正在做医学影像数据增强、小样本分割或者想了解最优传输如何落地的人。
1. OTLesMix 到底在解决什么问题
1.1 病灶样本不足不是简单缺图片
很多医学影像数据集里,病灶区域可能只占整个切片的 1% 都不到。以肝脏肿瘤分割为例,医生标注出来的病灶区域往往只有几个小连通域,大部分切片都是正常肝实质、血管和胆管。模型在训练时,如果按整图输入,网络会慢慢偏向“啥也不分”也能得到很高的分类准确率,因为背景实在太大了。即使使用类别权重,真正提供给分割网络的形态学习信号仍然有限,尤其是小病灶,一两个像素的误差就会被 Dice 指标放大得很厉害。
这时候常见做法是加旋转、翻转、缩放、弹性形变。这些增强确实简单有效,但它们本质上是在已有病灶上做几何变换,不能创造新的形状和新的位置。举个例子,如果数据集里的肺结节大多是直径 8mm 到 15mm 的类圆形病灶,你把它们旋转 30 度,得到的仍然是类圆形病灶。模型并没有见过更细长、分叶状、毛刺更明显的结节,泛化能力自然上不去。
OTLesMix 想解决的问题,是在给定有限真实病灶掩膜的情况下,自动生成一批形状更多样、位置更合理的合成病灶。它不是随机画一个椭圆贴上去,而是先学习真实病灶的分布规律,再在这个分布附近生成新样本。这样生成的病灶虽然不是来自真实病例,但在几何特征、边界形态、强度分布各方面都更接近能用于训练的正样本。
1.2 合成病灶要同时控制形状、位置和上下文
合成病灶看起来只要“画一块病变区域”就行,实际落地有三个难点。
第一个是形状多样性。真实病灶边缘不规则,内部信号也可能不均匀。简单的二阶平滑插值很容易把边界磨圆,让生成结果看起来像气泡,不像病灶。这里需要一种能在保持边界连续性的同时,允许形状变化的生成方式。
第二个是位置多样性。病灶不是均匀出现在图像任意位置。肝脏病灶基本在肝实质内,肺结节基本在肺野内,脑肿瘤基本在脑实质内,而且往往靠近特定组织结构。如果合成结果被随便放到器官外、骨头里、背景空气里,模型就会学到错误的空间先验。
第三个是上下文一致性。病灶和周围组织的关系很关键。CT 图像里一个低密度肝转移灶周围的肝实质灰度、血管走向、边缘对比度,都是医生判断的重要依据。如果只是把病灶掩膜硬贴上去,边界会出现明显断裂,模型可能学到的是“边界伪影对应病灶”,而不是真实病理特征。
OTLesMix 对这三个难点的处理方式,是把病灶生成看作一个分布层面的最优传输问题。先通过 Wasserstein Barycenter 从一组真实病灶中学习代表性模式,再用 Optimal Transport Map 把合成病灶从源区域移动到目标位置。这样生成结果就不是简单的像素插值,而是带着真实边界特征的结构迁移。
2. Wasserstein 重心与最优传输映射:合成病灶的数学底盘
2.1 用距离度量定义病灶之间的变化
最优传输最经典的理解是搬土问题:假设有一堆土,要从一个分布移动成另一个分布,搬运单位土方有成本,我们要找成本最低的搬运方案。Wasserstein 距离就是这种最优成本。放在病灶生成里,可以把一个病灶看成一片携带着强度信息的“土堆”,另一个病灶看成目标“土堆”,两者之间的 Wasserstein 距离描述了把一个病灶改造成另一个病灶需要的最小几何变化量。
相比之下,像素级 L2 距离对错位非常敏感:两个病灶明明形状接近,只是平移了一个像素,L2 距离就可能很大;而 Wasserstein 距离会认为这个变化很小,因为只需要把土往旁边挪一点。病灶形态分析里经常需要这种对轻微位移不敏感、对整体分布敏感的度量方式。
OTLesMix 使用这个度量的目的,是想在“真实病灶集合”和“目标位置区域”之间建立一个可计算的桥梁。有了距离,就可以定义重心,也可以定义映射。
2.2 重心求解:从一堆真实病灶里求一个模式
Wasserstein Barycenter 可以理解为一组分布在 Wasserstein 度量下的“平均分布”。对病灶来说,假设你有 N 个真实病灶 patch,每个 patch 的掩膜和一个强度剖面,重心就是一个新的病灶分布,它是这 N 个病灶在搬土距离下的加权折中。
这个重心不是简单地把 N 张图叠加取平均,那样边界会严重模糊。Wasserstein 重心会尽量保留各个样本里的结构,同时让整体形状处于它们的中间位置。你可以把它理解成“在形状和位置的几何约束下,综合出一类代表性病灶”。
实际计算时,最常见的做法是对每个 patch 做归一化,把它看成概率分布,然后求解带熵正则化的重心问题。熵正则化会引入一个参数 reg,通常是 Sinkhorn 散度里的正则系数。reg 越大,重心越平滑,边界越干净但细节可能丢失;reg 越小,重心越锐利,但计算越容易不稳定。这个参数需要根据图像尺寸和病灶尺度去调,不能直接套默认值。
2.3 最优传输映射:把合成结果放到指定位置
得到重心后,合成病灶还必须放到具体的目标区域里。这里就要用到 Optimal Transport Map。给定源分布和目标分布,最优传输映射会把源分布中的每个质量单元送到目标分布中合适的位置,并让总运输代价最小。
在 OTLesMix 的思路里,源分布可以是一个真实病灶 patch 或重心结果,目标分布可以是期望病灶出现的位置区域,也可以是一张目标图像上的局部上下文特征。通过传输映射,合成病灶会被“搬”到目标位置,同时保留它本身的结构特征。
这样做的好处是位置和形状可以解耦:形状由多个真实病灶的重心和传输过程控制,位置由目标分布的约束控制。换句话说,你可以保持一个类圆形病灶的边界特征,通过改变目标分布把它放到不同切片的不同位置;也可以让同一位置上的病灶呈现出更细长或更不规则的形态。
这一块不需要自己手写求解器。Python 生态里有 POT 库提供了 Wasserstein 距离、Sinkhorn、重心求解、EMD 等多种函数,能覆盖大部分原型验证。但你要理解每个函数输入的是矩阵还是概率向量,因为病灶数据往往是三维数组,直接喂给 OT 库里很容易踩维度坑。
3. 从环境准备到第一次把病灶合成出来
3.1 环境与依赖
先明确一点:OTLesMix 不是一个开箱即用的软件包,更接近一个算法思路。落地时你需要自己组装数据和流程。下面是一个比较常见的最小依赖组合,适合先跑通思路:
- Python 3.9 或更高版本
- PyTorch,如果后续要与分割网络联动
- NumPy、SciPy,做数组运算和几何变换
- scikit-image,做连通域分析和图像形态学处理
- SimpleITK 或 nibabel,读取 NIfTI、DICOM 等医学图像
- POT,处理 Wasserstein 重心和最优传输
- MONAI,可选,方便做医学图像预处理和增强
建议用虚拟环境,避免依赖冲突。一条可行的安装命令大致是这样:
conda create -n otlesmix python=3.9 -y conda activate otlesmix pip install torch numpy scipy scikit-image pot SimpleITK nibabel如果不需要训练网络,只做数据生成,可以暂时不装 PyTorch。但后面如果要验证增强效果,还是要有一个分割网络,所以建议直接装。
3.2 最小输入格式
OTLesMix 的输入最好是图像和掩膜成对出现。图像可以是 CT、MRI 或者普通病理切片,掩膜是二值标签或整数标签,表示病灶区域。建议先跑 2D 切片,因为 2D 的 patch 更小,OT 计算更快,可视化也直观;3D 体积数据可以等 2D 跑通后再扩展。
数据预处理要做两件最关键的事:
- 把图像和掩膜重采样到相同的体素间距或相同的分辨率,避免图像和掩膜错位。
- 把图像强度归一化到 0 到 1 或者 z-score 标准化,防止不同病例对比度差异干扰传输计算。
推荐用一张 CSV 管理样本列表,至少包含下面几列:
patient_id, image_path, mask_path, label, split LIDC-001, /data/LIDC/001.nii.gz, /data/LIDC/001_mask.nii.gz, 1, train LIDC-002, /data/LIDC/002.nii.gz, /data/LIDC/002_mask.nii.gz, 1, train这样后续批量生成时,不需要在每个脚本里硬编码路径,也方便做数据切分。
3.3 最小流程:加载、求解、合成、贴回
先给一个不绑定具体论文实现的通用流程伪代码,目的是帮助你理解数据流,而不是照抄就能得到论文效果。
# 伪代码:验证 OTLesMix 思路的最小流程 import numpy as np import SimpleITK as sitk from scipy import ndimage import ot def load_image_mask(image_path, mask_path): image = sitk.GetArrayFromImage(sitk.ReadImage(image_path)) mask = sitk.GetArrayFromImage(sitk.ReadImage(mask_path)) return image.astype(np.float32), mask.astype(np.float32) def extract_lesion_patches(image, mask, patch_size=64): # 找到掩膜连通域,计算质心和包围盒 labeled = ndimage.label(mask)[0] patches = [] for region_id in range(1, labeled.max() + 1): coords = np.argwhere(labeled == region_id) center = coords.mean(axis=0).astype(int) half = patch_size // 2 patch_coords = [slice(c - half, c + half) for c in center] patch_img = image[tuple(patch_coords)] patch_mask = (labeled[tuple(patch_coords)] == region_id).astype(np.float32) patches.append((patch_img, patch_mask)) return patches def barycenter_from_patches(patches, reg=0.01, num_iter=100): # 将每个 patch 展平成概率分布,然后求 Sinkhorn 重心 flatten_patches = [] for img, mask in patches: vec = (img * mask).reshape(-1) vec = vec / (vec.sum() + 1e-8) flatten_patches.append(vec) A = np.stack(flatten_patches, axis=1) bary_vec = ot.bregman.barycenter_sinkhorn(A, reg=reg, numItermax=num_iter) return bary_vec.reshape(patches[0][0].shape) def apply_transport_to_target(lesion, target_distribution): # 构造源分布和目标分布,调用 ot.emd 或 ot.sinkhorn 求映射 # 这里省略具体实现,核心是把 lesion 变成与 target 匹配的分布 transported = lesion return transported # 主流程 image, mask = load_image_mask("image.nii.gz", "mask.nii.gz") patches = extract_lesion_patches(image, mask, patch_size=64) bary = barycenter_from_patches(patches) synthetic = apply_transport_to_target(bary, target_distribution)这段代码只是流程示意,真正落地时你还需要处理多个细节:图像和掩膜的包围盒能不能超出边界,质心坐标取整后是否越界,多个病灶 patch 大小是否需要统一,目标分布从哪里来。我建议先跑一次最小样例,把输入输出打出来,确认图像、掩膜、patches 的维度是一致的,再进入后面的合成和贴回。
生成结果至少保存三样东西:合成图像、合成掩膜、合成图像叠加掩膜的可视化图。如果掩膜在贴回后没有和图像对齐,或者病灶位置落到了目标区域之外,可视化能一眼看出来。
4. 关键参数与质量判断标准
4.1 核心参数怎么调
OTLesMix 里有几个参数直接影响生成质量和计算成本。下面这张表可以作为调试起点,但实际参数必须结合你的图像尺寸、病灶大小和标注质量来定。
| 参数 | 作用 | 建议起点 |
|---|---|---|
| patch_size | 病灶周围上下文范围 | 典型病灶直径的 1.5 到 2 倍 |
| reg | Sinkhorn 熵正则化系数 | 0.01 到 0.1 |
| num_iter | 重心迭代次数 | 10 到 100 |
| batch_size | 每次参与重心计算的病灶数 | 8 到 32 |
| resolution | 是否重采样到各向同性体素 | 2D 可以先原分辨率,3D 先降采样 |
| transport metric | 构造代价矩阵用的距离类型 | 空间距离或强度距离 |
patch_size 虽然叫 patch,但它决定的不只是裁剪尺寸,而是整个传输计算的规模。patch 越大,传输矩阵越大,POT 内存占用越高。如果病灶直径只有 10 个像素,patch_size 开到 256 会让大部分区域都是背景,最优传输会把大量质量搬运到背景上,反而削弱病灶结构。
reg 值的调试逻辑是:图像噪声明显时,可以用稍大的 reg 平滑结果;如果生成病灶边缘毛刺明显,同时计算稳定,可以尝试更小的 reg。但 reg 太小会接近精确 EMD,对内存和迭代次数都非常敏感,不是越小越好。
4.2 生成结果怎么验证
别只看合成图是否“像”,要从三个层面验证。
首先是形状分布。把真实病灶和合成病灶的面积、周长、圆形度、长短轴比分别统计出来,画分布图,看两者是否接近。理想情况下,合成病灶的数量比真实病灶多,但分布不能偏离太多。如果合成病灶整体偏大或过于圆形,说明重心计算或后续形变过程压扁了边界多样性。
然后是位置分布。计算每个病灶的质心坐标,以及质心到目标器官掩膜边界的距离。比如目标区域是肝脏,合成病灶应该落在肝实质内。如果很多合成病灶落在肠道或腹壁,说明目标分布约束不够,需要在传输映射里加入解剖位置约束。
最后是下游任务验证。这是最实用的一步:拿固定训练集训练一个分割网络,比较三种策略——不做增强、普通几何增强、OTLesMix 增强。如果 OTLesMix 增强后的模型在真实验证集上 Dice 反而下降,说明合成数据离真实分布太远,可能需要降低合成样本比例或优化上下文融合。
4.3 一个推荐的小实验协议
我会先在训练集里随机抽 20 到 30 个真实病灶,拆成一个小的留出集,专门用来做质量评估。然后把 OTLesMix 生成的样本和真实样本混在一起,按不同比例训练分割模型,比如 1:0、1:1、1:3、0:1。每组固定训练轮数、优化器、学习率和随机种子。最后看验证集的 Dice 和 HD95。
这个小实验规模不大,跑完大概需要几小时,但能快速告诉你两个关键信息:合成数据到底有没有帮助,以及合成样本比例控制在多少最合适。不要一上来就生成几千张图,那样调参成本太高。
5. 批量增强与生产化落地
5.1 先单样本,再批量
很多人在跑通一个样例后,立刻开脚本批量生成几百个病灶,结果发现进度条卡在某一张图上,或者某些病例输出空掩膜。问题往往不是算法不行,而是没有把单样本流程做稳。
我建议先只处理一个病例,生成 5 到 10 个合成病灶,把以下信息全部打印出来:源图像路径、掩膜路径、病灶数量、每个 patch 的尺寸、重心计算耗时、生成结果保存路径。确认没有异常后,再扩展成完整的 CSV 列表。
批量时也不要开最大并发。OT 计算本身对内存很敏感,尤其是 3D 数据,并发一高很容易把内存打满。更稳妥的方式是单进程逐条处理,或者限制并发数为 2 到 4,同时监控 CPU 和内存使用。
5.2 输出目录、日志和失败重试
批量生成需要一套清晰的输出组织方式。下面是一个比较实用的目录结构:
outputs/ images/ masks/ overlays/ logs/ manifest.csvmanifest.csv 每一行记录一条合成记录,至少包含:
source_image, source_mask, synthetic_image, synthetic_mask, lesion_id, seed, reg, status这样以后想复现某个合成结果,可以直接根据 seed、reg 和 source_image 重新生成,不需要翻原始日志。
失败重试要单独考虑。批量时如果有一条数据读图失败,不要让整个任务中断。正确做法是记录失败原因,把这条数据放进 retry 队列,全部跑完后批量重试。重试仍然失败的数据,单独写一个 failed.csv,人工检查是路径问题、标注问题还是格式问题。
5.3 和 MixUp、CutMix、GAN、扩散模型怎么选
OTLesMix 不是唯一能做数据增强的方法,每种方法侧重点不一样。
| 方法 | 核心思路 | 优点 | 缺点 |
|---|---|---|---|
| MixUp | 图像和标签线性插值 | 实现简单,适合分类 | 对分割边界不友好 |
| CutMix | 把一个区域直接贴到另一张图 | 简单直接 | 边界和上下文不连续 |
| GAN | 对抗训练生成图像 | 视觉质量高 | 训练不稳定,配对掩膜难 |
| 扩散模型 | 逐步去噪生成图像 | 质量高,多样性强 | 采样慢,计算成本高 |
| OTLesMix | 最优传输加重心 | 可解释,不需要判别器,适合小样本 | 分布构造复杂,高分辨率计算成本高 |
如果只是做分类任务的图像增强,MixUp 和 CutMix 性价比最高,几行代码就能实现。如果做病灶分割,已经有精细掩膜,那么 CutMix 容易露出贴图痕迹,GAN 又需要额外训练判别器,OTLesMix 的优势就更明显。它不依赖额外判别器,输入输出都是图像加掩膜配对,和分割训练流程天然匹配。
但也要注意,OTLesMix 并不是生成任意新图像的工具,它的强项是在已有病灶基础上做形态和位置的扩展。你拿不到比真实训练集更有用的病理结构信息,它只负责把已有信息组合得更多样。
6. 常见报错与排查顺序
6.1 先看数据,再看依赖,最后才调参数
遇到问题不要第一时间改 reg,也不要重新写模型。我一般按这个顺序排查:先看现象,是报错、卡住、输出全黑还是输出质量差;再看数据,图像和掩膜是否对齐、路径是否正确、病灶掩膜是否二值化;然后看依赖版本,POT、SimpleITK、PyTorch 的接口有没有变化;最后才调整参数。
很多情况下,问题出在一个很小的前置环节。比如掩膜是 int16 类型,里面除了 0 和 1 之外还有 255 的标记值,归一化后病灶区域变成多个不同强度;再比如图像和掩膜一个用 DICOM 读取,一个用 NIfTI 读取,体素间距不一致导致错位。这些都不是算法问题,但会让后续所有步骤都变得奇怪。
6.2 典型异常现象与优先排查项
| 现象 | 优先排查 |
|---|---|
| 生成病灶全黑 | 输入 mask 是否二值化,patch 里是否包含病灶,归一化是否过强 |
| 边界毛刺明显 | reg 是否太小,patch 是否在原分辨率上直接计算,缺少形态学平滑 |
| 内存暴涨 | patch 是否过大,3D 是否直接全分辨率运行,Sinkhorn 迭代是否太多 |
| 生成结果几乎一样 | 重心占比过高,样本多样性不足,需要增加随机采样或调整合成比例 |
| 训练时 loss 震荡明显 | 合成数据比例过高,或者生成样本与真实分布偏移太大 |
| 读取文件报错 | 路径是否含特殊字符,图像和掩膜是否存在,SimpleITK 是否支持该格式 |
POT 里常见的一个坑是输入矩阵包含 NaN 或全零向量。病灶 patch 可能一整块都是背景,展平成概率分布后 sum 为 0,直接传入 barycenter 函数会报错。处理方式是在展平前先过滤掉病灶面积过小的 patch,或者给除零位置加上一个极小 epsilon。
6.3 优化方向和使用边界
如果计算资源有限,先降分辨率验证算法,再逐步提高。比如 2D patch 先用 64×64,在少量样本上跑通后,再扩大到 128×128。3D 数据可以先用各向同性重采样降低体素数,跑通后再回到原始分辨率。
如果希望生成结果更贴近解剖结构,可以考虑在传输映射中加入解剖位置约束。比如目标区域是肝脏,就提前准备一个肝脏掩膜,让病灶质心必须落在肝脏掩膜内,并且避免覆盖大血管或胆管。这种规则性约束比单纯依赖 OT 更可靠,因为 OT 不知道解剖学语义。
最后说一句边界:OTLesMix 不能替代真实标注,也不能解决标注质量本身的问题。如果原始病灶掩膜边缘标注得很差,传输之后仍然会继承错误边缘。它的价值是在有限真实样本基础上扩大形态和位置覆盖范围,让模型在遇到分布内变化时更稳定,而不是凭空创造全新的病变类型。
很多失败其实不是模型理论问题,而是前置流程没做干净:掩膜没有重采样到和图像一致,路径含特殊字符导致读取异常,Sinkhorn 正则化设得太小导致迭代不收敛。把单条样本跑稳,再考虑批量生成和调参,这才是把 OTLesMix 落到自己数据集上的正确顺序。