前言
这次我们来看一个医学图像 AI 领域比较有代表性的方法:OTLesMix。它的核心不是做一个大模型,而是解决一个很实际的问题——病灶太少了,模型学不够。在医学影像中,带标签的病灶数据往往稀缺、形态多样、位置分散,而模型要学到的恰好又是“形态变化大、位置分布广”的病灶特征。OTLesMix 的思路很直接:既然真实病灶少,就用数学工具按形态和位置“合成”出更多合理的病灶。它不靠 GAN,也不靠扩散模型,而是引入最优传输里的Wasserstein Barycenter和Optimal Transport Map来生成病灶。
这篇文章会围绕四个问题展开:
- OTLesMix 到底做了什么,和 CutMix、MixUp 这些常见增强方法有什么区别。
- Wasserstein Barycenter 和 Optimal Transport Map 在病灶生成中分别承担什么角色。
- 如果要从代码层面复现或验证这套方法,环境、数据、训练流程怎么搭。
- 实际使用时有哪些坑,显存、数据格式、评估指标分别怎么看。
无论你是做医学图像分割、检测,还是做数据增强方向的研究,这篇文章都值得往下看。
1. 核心能力速览
| 能力项 | 说明 |
|---|---|
| 项目类型 | 医学图像合成病灶生成方法,研究方向/算法级实现 |
| 核心方法 | Wasserstein Barycenter 求病灶形态重心,Optimal Transport Map 实现病灶形态与位置迁移 |
| 主要功能 | 从已有病灶生成形态更多样、位置更分散的合成病灶,用于增强分割/检测训练集 |
| 适用模态 | 常见于 CT、MRI 等 3D 医学影像,也可扩展到 2D 切片 |
| 与 GAN 类的区别 | 不需要对抗训练,依赖最优传输数学框架,合成过程更可控 |
| 需要哪些额外数据 | 病灶掩码(mask)或标注框,不需要额外的真实图像 |
| 硬件门槛 | 训练阶段建议 GPU,显存需按 2D/3D 和输入尺寸测试;推理阶段单张图像生成可尝试 CPU |
| 是否支持一键启动 | 取决于作者是否提供完整代码仓库,通常需要自行配置环境 |
| 是否支持 API | 论文方法一般不带 Web API,需要自己封装 |
| 是否支持批量任务 | 合成病灶过程天然支持批量处理,适合离线数据增强 |
| 适合场景 | 小样本病灶分割、不平衡数据集增强、域泛化研究 |
需要提醒的是:目前公开材料中没有作者提供的完整可用代码仓库时,下面所有代码都是通用模板,用来表达实现思路,不能直接当作项目脚本运行。实际部署要以具体开源实现为准。
2. 方法原理:Wasserstein Barycenter 和 Optimal Transport Map
2.1 为什么不能用普通的 CutMix / MixUp
在自然图像里,CutMix 的用法是:从一张图里裁一块放到另一张图上,标签也跟着混合。这个方法在分类任务里有效,但在医学病灶增强里会有问题:
- 病灶不是规则矩形,直接裁切会把健康组织也切进来。
- 病灶形态多变,矩形拼接生成的图像看起来不真实。
- 位置随意摆放可能违背人体结构约束,比如病灶出现在不该出现的解剖位置。
- 医学图像通常要求掩码精确到像素级,简单拼贴会产生明显边界伪影。
OTLesMix 的思路是把“形态”和“位置”分开处理:用最优传输的方法让病灶形状发生可控变形,再通过重心计算和映射把病灶放到新的合理位置。
2.2 最优传输:将病灶形态变化建模成映射问题
最优传输(Optimal Transport, OT)研究的是“如何以最小代价把一个分布搬运到另一个分布”。病灶形态本质上也可以看作一种二维或三维空间中的质量分布。比如病灶 A 和病灶 B,我们可以把 A 的像素密度分布通过一个最优传输映射变换到 B 的分布附近。
这里的关键是:最优传输映射天然带有几何意义。它给出的不只是“哪些点变到哪些点”,还保留了空间的连续性和拓扑结构。因此用它来变形病灶,不会把病灶撕裂成不连贯的区域。
在 OTLesMix 里,Optimal Transport Map 承担的任务主要是:
- 给定源病灶和目标病灶,找到两者之间像素位置对应关系;
- 根据对应关系对源病灶进行插值或变形;
- 产生从“源形态”过渡到“目标形态”的中间形态。
这样就不需要训练一个生成模型,而是直接通过数学计算完成形态插值。
2.3 Wasserstein Barycenter:多病灶形态求重心
Wasserstein 距离也叫推土机距离(Earth Mover's Distance),描述的是把一种分布变成另一种分布所需的最小搬运成本。
Wasserstein Barycenter 则是多个这种分布的“重心”。给定 N 个病灶掩码,计算它们在 Wasserstein 意义上的平均值,得到一个能代表这群病灶“共同形态特征”的掩码。
这个重心有什么用?举个例子:
- 输入 3 个肝肿瘤掩码:一个偏圆、一个偏长条、一个形态不规则;
- 它们的 Wasserstein Barycenter 可能是一个覆盖三者形态特征的形态;
- 通过在各个病灶与重心之间做插值参数 t 调控,就能得到一系列介于两者之间的新病灶形态;
- t 越接近 0,越接近原始病灶;越接近 1,越接近重心。
换句话说,Barycenter 的作用是提供一个“形态锚点”,配合最优传输映射,让合成病灶既保有原来的纹理特征,又能产生新的合理形态。
2.4 OTLesMix 与 GAN、扩散模型的区别
| 对比项 | OTLesMix/OT 方式 | GAN | 扩散模型 |
|---|---|---|---|
| 是否需要对抗训练 | 否 | 是 | 否 |
| 实现难度 | 数学计算,相对轻量 | 训练不稳定,调参难 | 训练和采样成本高 |
| 合成可控性 | 有明确的形态插值和位置映射,较可控 | 通常控制粒度较粗 | 需要额外引导才能精确控制 |
| 数据需求 | 有病灶掩码即可 | 需要大量训练数据 | 需要大量训练数据 |
| 计算成本 | 低到中 | 中到高 | 高 |
| 适合场景 | 小样本增强、数据扩充 | 大规模图像生成 | 高质量生成但资源消耗大 |
OTLesMix 的优势恰恰在于它对数据量和计算资源的要求比较低。它不是从头“画”一个病灶,而是基于已有病灶做形态和位置的“插值迁移”,因此在小样本的医学影像场景下更实用。
3. 适用场景与使用边界
3.1 适合谁用
- 做病灶分割的研究者:训练集里病灶数量不足,想扩展训练数据。
- 做小样本医学图像分类的人:利用合成病灶增强稀有类别。
- 做域泛化研究的人:希望模型见过更多“形态变化”的病灶,提高泛化性。
- 不想引入 GAN/扩散模型复杂训练流程的人:用最优传输做数据增强,流程更轻。
3.2 不适合什么场景
- 如果已经拥有大量真实病灶标注数据,OTLesMix 的价值主要体现在数据多样性上,不能替代真实数据。
- 如果病灶没有明显形态边界(例如弥漫性病变),掩码本身难以定义,最优传输的效果会受限制。
- 如果下游任务要求合成图像必须和临床完全一致,这类合成增强方法的可靠性需要额外验证。
3.3 使用边界与合规提醒
医学图像合成病灶存在明确的伦理和合规要求:
- 合成数据只能用于研究目的的模型训练验证,不能直接作为临床诊断依据。
- 涉及真实患者数据时,必须遵守医院伦理审批、数据脱敏、患者隐私保护要求。
- 生成的病灶不能用于伪造病例、学术造假或误导诊断。
- 如果要发布合成图像数据集,必须明确说明合成方式和适用边界。
4. 环境准备与前置条件
4.1 基础环境清单
复现或参考 OTLesMix 这类方法,环境上至少需要满足以下条件:
| 依赖项 | 说明 |
|---|---|
| 操作系统 | Linux 优先,Windows/macOS 也能跑但要注意依赖兼容 |
| Python | 推荐 3.8 到 3.11,具体看 PyTorch 版本 |
| PyTorch | 推荐 CUDA 版,2.0 以上较稳定 |
| CUDA | 推荐 11.8 或 12.x,看驱动版本 |
| 医学图像库 | SimpleITK、NiBabel、MONAI,用于读 nii.gz 等医学格式 |
| 最优传输计算库 | 推荐 POT(Python Optimal Transport),这是最常用的 OT 库 |
| GPU | 有 NVIDIA GPU 更好;CPU 可以跑小尺寸实验但较慢 |
4.2 安装 POT 库
OTLesMix 的核心数学计算主要依赖最优传输工具包。最常用的库是POT,安装命令:
pip install POTPOT 库提供了 Wasserstein 距离计算、Sinkhorn 近似、Barycenter 计算等方法,是复现 OTLesMix 最快的起点。
4.3 安装医学图像处理依赖
如果你是处理 3D 医学图像(如 .nii.gz 格式),推荐安装:
pip install monai==1.3.0 pip install simpleitk nibabel pip install numpy scipy scikit-image pip install matplotlib tqdmMONAI 是为了方便做医学图像预处理和数据加载;SimpleITK 和 NiBabel 负责格式读写。如果只是做 2D 实验,SimpleITK 配合 OpenCV 也可以。
4.4 显存和磁盘空间
显存需求取决于你处理的是 2D 切片还是 3D 体数据,以及输入分辨率。这里给出一个经验区间,但最终占用必须按你的实际设置测试:
| 输入类型 | 推荐显存 | 说明 |
|---|---|---|
| 2D 切片 | 6G 到 12G | 批量小的话 6G 可以跑 |
| 3D 小体块裁剪 | 12G 到 24G | 裁剪区域越小越省显存 |
| 3D 全尺寸 | 24G 以上 | 全分辨率训练较吃资源 |
磁盘空间主要看数据量。医学图像一个样本通常几十 MB 到几百 MB,如果做大批量合成增强,建议预留几十 GB 空间存放生成结果。
5. 数据准备与病灶掩码处理
OTLesMix 方法对数据格式有一定要求。核心输入是“病灶掩码”和“原始图像”。这里的掩码指的是标注的病灶区域,可以是二值掩码,也可以是概率图。
5.1 数据目录结构建议
data/ ├── images/ │ ├── patient_001.nii.gz │ ├── patient_002.nii.gz │ └── ... ├── labels/ │ ├── patient_001.nii.gz │ ├── patient_002.nii.gz │ └── ... └── synthetic/ ├── images/ └── labels/5.2 读取病灶掩码示例
下面是一段基于 MONAI 的读图模板代码:
import torch from monai.data import ImageReader reader = ImageReader() # 读取图像和标签 image, image_meta = reader.read("data/images/patient_001.nii.gz") label, label_meta = reader.read("data/labels/patient_001.nii.gz") image_tensor = torch.tensor(image).float() label_tensor = torch.tensor(label).float() # 二值化,确保掩码取值为 0 或 1 label_binary = (label_tensor > 0.5).float() print("image shape:", image_tensor.shape) print("label shape:", label_binary.shape) print("病灶像素比例:", label_binary.sum() / label_binary.numel())5.3 提取病灶连通域
实际场景里,一个病人可能只有一处病灶,也可能有多处。每个病灶的掩码应该单独提取出来作为 OT 计算的输入单位。常见做法是用连通域分析:
from scipy.ndimage import label # 对二值掩码做连通域标记 labeled_array, num_features = label(label_binary.numpy()) lesion_masks = [] for i in range(1, num_features + 1): lesion_mask = (labeled_array == i).astype(float) lesion_masks.append(lesion_mask) print("检测到病灶数量:", len(lesion_masks))这一步很重要,因为 Wasserstein Barycenter 计算的输入是“一个独立的病灶形态”,而不是一整幅图。如果把多个病灶混在一起算重心,形态会失真。
5.4 归一化和重采样
不同病人的 CT 图像,像素间距和灰度范围可能不一致。建议在计算前做两件事:
- 重采样到统一的 spacing,例如
1.0 x 1.0 x 1.0mm; - 灰度裁剪到 [0,1] 范围,或使用窗宽窗位归一化。
6. OTLesMix 核心计算流程
OTLesMix 的整体流程可以拆成四个阶段:
- 从训练集中提取病灶掩码集合。
- 计算一组病灶的 Wasserstein Barycenter。
- 基于最优传输映射,将源病灶向目标位置或目标形态迁移。
- 把生成的合成病灶贴回健康组织的对应位置。
6.1 计算 Wasserstein 距离
先看如何用 POT 计算两个病灶掩码之间的 Wasserstein 距离。以下代码针对 2D 掩码:
import numpy as np import ot def mask_to_points(mask, num_samples=1000): """ 将二值掩码转换为点集。 只采样掩码内的坐标点。 """ coords = np.argwhere(mask > 0.5) if len(coords) == 0: return np.zeros((1, 2)) if len(coords) > num_samples: idx = np.random.choice(len(coords), num_samples, replace=False) coords = coords[idx] return coords.astype(float) # 假设两个 2D 病灶掩码 mask_a = np.zeros((128, 128)) mask_b = np.zeros((128, 128)) mask_a[40:80, 30:70] = 1 mask_b[60:100, 50:90] = 1 points_a = mask_to_points(mask_a) points_b = mask_to_points(mask_b) # 计算距离矩阵 M = np.sum( (points_a[:, None, :] - points_b[None, :, :]) ** 2, axis=-1 ) # 归一化权重 a = np.ones(points_a.shape[0]) / points_a.shape[0] b = np.ones(points_b.shape[0]) / points_b.shape[0] # 计算 Wasserstein 距离 wd = ot.emd2(a, b, M, numItermax=200000) print("Wasserstein Distance:", wd)这段代码的核心是:把病灶掩码看成点的分布,计算源点云到目标点云的最小搬运距离。距离越小,说明两个病灶形态越接近。
6.2 计算 Wasserstein Barycenter
POT 库提供了计算 Wasserstein Barycenter 的接口。对一组病灶掩码,可以这样近似计算重心:
import numpy as np import ot def compute_barycenter(mask_list, num_bins_x=64, num_bins_y=64, reg=0.01): """ 在批量网格上近似计算 Wasserstein Barycenter。 mask_list: 列表,每个元素是 (H, W) 的二值掩码。 """ ref_h, ref_w = mask_list[0].shape grid_x = np.linspace(0, 1, num_bins_x) grid_y = np.linspace(0, 1, num_bins_y) xx, yy = np.meshgrid(grid_x, grid_y) grid = np.stack([xx.ravel(), yy.ravel()], axis=1) M = np.sum( (grid[:, None, :] - grid[None, :, :]) ** 2, axis=-1 ) M /= M.max() distributions = [] for mask in mask_list: # 把掩码缩放到网格大小 hist = _resample_mask(mask, num_bins_y, num_bins_x) hist = hist.ravel() if hist.sum() == 0: hist = np.ones_like(hist) / len(hist) else: hist = hist / hist.sum() distributions.append(hist) weights = np.ones(len(distributions)) / len(distributions) bary = ot.bregman.barycenter( A=np.array(distributions).T, M=M, reg=reg, weights=weights ) return bary.reshape(num_bins_y, num_bins_x) def _resample_mask(mask, h, w): from scipy.ndimage import zoom zoom_y = h / mask.shape[0] zoom_x = w / mask.shape[1] resized = zoom(mask, (zoom_y, zoom_x), order=1) return (resized > 0.5).astype(float)需要留意的是,真正的 Wasserstein Barycenter 在离散点云上计算十分耗时。POT库提供了多种近似算法,包括 Sinkhorn 正则化方法。上面代码中加了正则项reg,就是 Sinkhorn 近似。实际使用时需要根据病灶大小调整网格数和正则强度。
6.3 最优传输映射与形态插值
得到重心后,OTLesMix 会通过最优传输映射(Optimal Transport Map)在两个病灶之间做形态插值。简单说:
- 把源病灶的点用最优传输映射搬运到目标位置;
- 插值参数 t 控制搬运的完成度;
- t 步进扫描可以得到一系列中间形态。
def interpolate_lesion(mask_source, mask_target, t=0.5): """ 在源病灶和目标病灶之间进行最优传输插值。 简化实现:基于点集插值。 """ points_src = mask_to_points(mask_source, num_samples=500) points_tgt = mask_to_points(mask_target, num_samples=500) # 计算代价矩阵 M = np.sum( (points_src[:, None, :] - points_tgt[None, :, :]) ** 2, axis=-1 ) # 求解最优传输规划 a = np.ones(points_src.shape[0]) / points_src.shape[0] b = np.ones(points_tgt.shape[0]) / points_tgt.shape[0] # 使用 Sinkhorn 求解 reg = 0.01 G = ot.sinkhorn(a, b, M, reg) # 插值映射 interp_points = (1 - t) * points_src + t * (G @ points_tgt) # 用插值后的点重建掩码 mask_result = np.zeros_like(mask_source) for (x, y) in interp_points.astype(int): if 0 <= x < mask_result.shape[0] and 0 <= y < mask_result.shape[1]: mask_result[x, y] = 1 from scipy.ndimage import binary_closing mask_result = binary_closing(mask_result, iterations=2).astype(float) return mask_result6.4 位置迁移合成
OTLesMix 的另一个重要能力是让病灶“换位置”。这背后的思路是:把病灶的最优传输映射推广到空间位置层面,在健康组织上选择合理的解剖位置插入合成病灶。
要实现这一点,需要先定义“可插入区域”。不成熟的暴力做法是随机选一个点然后贴上去。更好的做法是限制到目标器官区域内部。具体实现可能涉及解剖先验知识,例如肝病灶只能出现在肝脏分割区域内。
7. 训练增强流程示例
OTLesMix 在训练中的角色是数据增强模块。一个典型的分割训练 pipeline 应该这样组织:
- 每个 epoch 开始前,随机抽取一组标注病灶;
- 计算重心、做形态插值、生成新病灶;
- 将新病灶嵌入到健康图像中;
- 合成图像和合成掩码进入分割模型完成一次前向传播和反向传播。
7.1 合成病灶嵌入健康图像
下面是一个简化的合成嵌入流程模板:
def paste_lesion(image, healthy_mask, lesion_mask, position): """ 将合成病灶贴到健康区域的图像上。 image: 原始图像 healthy_mask: 健康组织区域掩码 lesion_mask: 合成病灶掩码 position: 插入位置 (x, y) """ x, y = position h, w = lesion_mask.shape # 检查越界 if x + h > image.shape[0] or y + w > image.shape[1]: return image, None # 确保插入位置在健康区域内 region = healthy_mask[x:x+h, y:y+w] if region.sum() < 0.5 * h * w: return image, None # 生成合成掩码 new_label = np.zeros_like(image) new_label[x:x+h, y:y+w] = lesion_mask # 把病灶纹理覆盖到图像上 bg_mean = image[x:x+h, y:y+w][lesion_mask > 0.5].mean() lesion_texture = image[x:x+h, y:y+w] * (1 - lesion_mask) synthesized = lesion_texture + lesion_mask * bg_mean new_image = image.copy() new_image[x:x+h, y:y+w] = np.where(lesion_mask > 0.5, synthesized, image[x:x+h, y:y+w]) return new_image, new_label这里采用了极简的纹理复制策略。实际论文中极大概率会根据病灶内部的灰度分布、纹理信息做更精细的处理,否则合成病灶和周围组织之间会有明显边界。真实实现通常还会加入高斯平滑、边界融合等步骤。
7.2 训练循环集成
在实际训练中,OTLesMix 模块可以做成一个独立的增强函数,挂在 dataloader 之前:
def otlesmix_augment(image, label, lesion_pool): """ 在训练时从病灶池中合成新病灶并嵌入当前图像。 """ if len(lesion_pool) < 2 or np.random.random() < 0.5: return image, label # 随机选两个病灶 idx1, idx2 = np.random.choice(len(lesion_pool), 2, replace=False) lesion_a = lesion_pool[idx1] lesion_b = lesion_pool[idx2] # 随机插值系数 t = np.random.uniform(0.2, 0.8) synthetic_mask = interpolate_lesion(lesion_a, lesion_b, t) # 嵌入到当前图中 h_w = synthetic_mask.shape pos_x = np.random.randint(0, image.shape[0] - h_w[0]) pos_y = np.random.randint(0, image.shape[1] - h_w[1]) new_image, new_label = paste_lesion( image, label == 0, synthetic_mask, (pos_x, pos_y) ) if new_label is None: return image, label return new_image, new_label这个模块的设计思路是“能用就用,不行就退回原图”,保证增强过程不会破坏训练稳定性。
8. 效果验证与评估指标
合成病灶生成方法好不好,不能只看生成图像“像不像”,更关键的是它能不能提升下游任务的性能。建议从三个层次验证。
8.1 视觉质量验证
- 检查合成病灶形态是否连续,是否有空洞、撕裂;
- 检查合成病灶边界是否与周围组织过渡自然;
- 对比原始病灶直方图和合成病灶直方图,观察灰度分布差异;
- 使用分割模型对其推理,观察预测掩码是否稳定。
8.2 下游任务性能验证
以病灶分割为例:
| 验证方式 | 说明 |
|---|---|
| 原始训练集训练 | 作为基线 |
| 原始训练集 + 随机 CutMix 增强 | 作为对比方法 |
| 原始训练集 + OTLesMix 风格增强 | 实验组 |
| 在独立测试集上比较 Dice、IoU、HD95 | 衡量分割精度 |
8.3 多样性验证
多样性是合成数据的一个关键指标。如果生成的病灶全是一个模子,那么增强效果会大打折扣。
可以统计一组生成病灶的:
- Wasserstein 距离两两分布;
- 形态特征分布(面积、长轴短轴比);
- 位置坐标分布。
如果这些统计量相对分散,说明合成病灶确实覆盖了更大的形态和位置空间。
8.4 计算代码示例
import numpy as np def compute_dice(pred, gt): """ 计算 Dice 系数。 """ intersection = np.sum(pred * gt) union = np.sum(pred) + np.sum(gt) if union == 0: return 1.0 return 2.0 * intersection / union9. 资源占用与性能观察
9.1 显存占用观察方式
在训练过程中,关注显存占用的方式最直接的是使用nvidia-smi:
watch -n 1 nvidia-smi也可以使用 PyTorch 自带的内存统计:
import torch def print_memory_usage(): if torch.cuda.is_available(): allocated = torch.cuda.memory_allocated() / 1024**3 reserved = torch.cuda.memory_reserved() / 1024**3 print(f"显存分配: {allocated:.2f} GB") print(f"显存预留: {reserved:.2f} GB")9.2 影响资源占用的关键因素
- 病灶裁剪尺寸:OTLesMix 通常先裁出病灶区域计算,区域越大,计算量越大;
- 点云采样数量:
mask_to_points中的num_samples越大,OT 计算越慢,内存越高; - Sinkhorn 正则强度:正则值越小收敛越慢,计算越久;
- 网格分辨率:Barycenter 计算的网格数(比如 64x64 vs 256x256)直接影响显存和耗时;
- 批量大小:合成增强如果在线生成,批量大会显著增加显存开销。
9.3 CPU vs GPU
Wasserstein Barycenter 和最优传输映射的主体计算涉及矩阵运算。小尺寸(2D切片、网格数64x64)在 CPU 上就能跑,但速度较慢。3D体数据、网格复杂度更高时,建议用 GPU。POT 库底层支持部分 GPU 加速,但不是所有接口都有 CUDA 实现,实际使用时要查当前版本的 API 文档。
9.4 降低资源占用的建议
- 离线预处理阶段生成合成数据,不要把 OT 计算放到训练循环里实时跑;
- 对病灶做中心裁剪和归一化,统一到固定尺寸;
- 降低
num_samples和 Barycenter 网格分辨率; - 适当增大 Sinkhorn 正则项,收敛更快;
- 使用混合精度训练,减少显存占用。
10. 常见问题与排查方法
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
| 安装 POT 失败 | Python 版本不兼容或缺少编译器 | 查看 pip 安装日志 | 升级 Python 或使用 conda 安装 |
| 读取 nii.gz 失败 | 文件路径错误或格式损坏 | 用 SimpleITK 单独打开 | 检查数据格式,替换损坏文件 |
| Barycenter 计算耗时过长 | 网格过大或正则过小 | 打印每次计算耗时 | 降低网格分辨率,增大 reg |
| 计算出现 NaN | Sinkhorn 迭代发散或数据含 0 分布 | 加断点检查中间值 | 对 hist 加平滑,增加正则项 |
| 合成病灶掩码有空洞 | 点云重建掩码时分辨率不足 | 可视化中间掩码 | 增加采样点数,做形态学闭合 |
| 嵌入位置出现重叠 | 插入位置未检查健康区域约束 | 可视化合成结果 | 增加健康区域判断逻辑 |
| 训练精度没有提升 | 合成病灶与真实病灶分布差距过大 | 统计合成数据和真实数据的分布 | 调整插值系数范围,结合原始数据混合训练 |
| 显存溢出 | 3D 计算或 batch 过大 | 观察 nvidia-smi | 降低 batch、裁剪尺寸、使用梯度累积 |
10.1 关于“效果不如预期”的排查思路
这是最容易遇到的问题。合成增强方法并不是“加上就一定涨点”。如果实验结果显示没有提升,可以从以下几个方向排查:
- 合成病灶是否在形态上和真实病灶差异太大,导致模型学到了错误的特征。
- 嵌入位置是否合理,如果病灶被放到组织边界处,模型学习会混乱。
- 合成数据比例是否合理,过高的合成比例可能降低模型对真实数据的敏感性。
- 插值系数 t 的分布是否过于集中在 0.5,缺少形态多样性。
11. 最佳实践与使用建议
11.1 数据层面
- 建立病灶池时,尽量覆盖不同形态、大小、位置的病灶。
- 对病灶做尺度归一化,避免大小差异过大影响 OT 计算稳定性。
- 对掩码做形态学后处理,去掉孤立点和零散小区域。
- 划分训练集、验证集、测试集时确保患者级分离,避免同一个患者出现在不同集合中。
11.2 实验设计层面
- 先离线生成一批合成数据,观察可视化和统计分布,确认合理后再进入训练。
- 基线实验、CutMix 对照、OTLesMix 对照必须用相同训练配置。
- 控制变量:调整插值 t、合成比例、嵌入位置策略时,一次只改一个变量。
- 每个实验至少重复两次,记录均值和方差。
11.3 工程化层面
- 合成阶段和训练阶段分离:合成数据先生成好存硬盘,训练时直接读取,不占训练时间。
- 批量生成时记录日志,包括输入病灶 ID、插值系数 t、位置坐标,方便复现。
- 对合成数据做质量过滤,例如计算合成掩码的面积是否在合理范围内。
- 编码时注意使用随机种子,保证实验可复现。
11.4 合规层面
- 使用真实医学图像时,必须确认数据授权范围。
- 合成病灶涉及患者隐私数据脱敏,不能直接暴露原始患者信息。
- 涉及人脸、器官、肿瘤等敏感医学内容时,不允许在未经授权的情况下公开合成样例。
12. 总结与下一步
OTLesMix 这类基于 Wasserstein Barycenter 和 Optimal Transport Map 的合成病灶生成方法,最大的价值在于:提供了一种不依赖 GAN 和扩散模型、计算相对可控、语义上可解释的数据增强思路。它把病灶生成拆分成了“形态插值”和“位置迁移”两步,每一步都有明确的数学定义,因此特别适合医学图像小样本场景。
如果你准备在自己的数据集上尝试,建议从这三步开始:
- 先准备一个病灶掩码池,统计它们的形态、大小分布。
- 用 POT 库跑通两个病灶之间的 Wasserstein 距离和 Sinkhorn 插值,确认能生成连续合理的中间形态。
- 再把合成病灶嵌入到健康图像中,按 8:1:1 的比例混合真实和合成数据,对比分割模型 Dice 指标是否有提升。
最容易踩的坑有两个:一是把 OT 计算直接塞进训练循环导致训练变慢;二是不控制合成病灶的解剖合理性,把病灶放在了错误位置。前者建议离线生成,后者建议加入区域约束。
如果你已经在用 GAN 或扩散模型做病灶增强,也可以尝试把 OTLesMix 当作前置筛选器,用 OT 方式生成多样化候选病灶,再用更重的生成模型做精修,两者互补而不是互斥。方法本身并不复杂,复杂的是如何把它嵌入到正确的数据流和实验设计里。建议先小规模跑通,再逐步扩大。