尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

双流多实例学习:从全切片图像到弱监督肿瘤检测与定位

双流多实例学习:从全切片图像到弱监督肿瘤检测与定位 拿到一张全切片图像whole slide imageWSI你手里唯一的监督信号是“这张片子有没有肿瘤”连具体位置都不告诉你——这时该怎么训练一个能同时给出诊断结果和病灶热图的模型这就是 multiple instance learning多实例学习MIL在计算病理领域最经典的应用场景。这篇论文笔记要拆解的《Dual-stream multiple instance learning networks for tumor detection in Whole Slide Image》核心就是把 MIL 的两条主流路线拧成一股绳一个支路学全局表示做分类另一个支路用高斯混合模型自动挑出最像肿瘤的关键 patch 做定位两个流共享特征、互相助力。文章按论文动机、双流架构、实验解读、代码复现、训练避坑几个部分展开适合正打算入坑弱监督病理图像分析、或者想在自己任务里复现“既能分类又能定位”机制的朋友参考。1. 论文动机与问题本质为什么一张 slide 只给一个标签也能训练1.1 WSI 分析的现实困境WSI 有多大一张典型的 40 倍病理切片扫描图分辨率通常能到 100000×100000 像素级别直接丢给卷积网络做像素级分类显存和算力都撑不住。常规做法是把 WSI 切成若干个小 patch比如 256×256 或者 512×512再用这些 patch 训练模型。但这里有个非常现实的问题如果要求医生逐像素标注肿瘤区域成本高到难以接受一个病理医生标注一张切片可能就要花几个小时而 slide 级别的“有/无肿瘤”标签在临床工作流里本来就存在几乎零成本获取。于是大家想到一个折中方案把一张 WSI 看成是一个“包”bag切出来的 patch 就是包里的“实例”instance。如果这张片子有肿瘤那包里至少有一个 patch 应该被判定为肿瘤如果这张片子没有肿瘤那包里所有 patch 都是正常组织。模型不需要知道哪个 patch 是肿瘤只需要在“包”的级别上做监督训练。这个设定天然就是多实例学习问题也是目前弱监督病理图像分析最主流的框架。1.2 两代 MIL 范式的分歧与死结MIL 在深度学习里大致分成两个流派理解它们的区别是读懂这篇论文的前提。第一个流派叫 embedding-based MIL思路是把包内所有实例的特征向量喂进一个排列不变的聚合函数比如 max pooling、mean pooling、attention-based pooling得到整个包的一个固定维度表示再接分类器输出包级标签。这种方法的优点是训练稳定、分类性能通常不错缺点是很不好做病灶定位。虽然 attention 机制能给出一个每个 patch 的注意力权重但那只是“注意力的相对大小”不是严格意义上的“肿瘤概率”在临床解释上不够直接。第二个流派叫 instance-based MIL思路是先对每个 patch 单独预测一个概率再用某种方式把实例级概率聚合成包级概率例如取最大值、取平均值、或者对 top-k 个实例做平均。最大池化版本是最原始的“至少一个阳性”假设的直接实现只要有一个 patch 被判定为阳性整个包就是阳性。这个流派的好处是天然可以输出每个 patch 的预测概率用来生成热图定位病灶位置坏处是 patch 级别预测噪声大容易把正常组织的某个局部纹理误判成肿瘤包级精度普遍不如 embedding-based。这两条路线看起来像是互斥的想要可解释性就牺牲精度想要精度就很难做像素级定位。这篇 Dual-stream MIL 论文做的恰恰不是二选一而是同时走两条路。一个支路在 embedding 空间做全局聚合另一个支路在实例空间做预测再用一个跨流共享机制把两边信息混起来最后用高斯混合模型动态筛选关键实例强调“先判断哪里可疑再根据可疑区域下结论”。2. 网络结构拆解双流怎么合并不冲突2.1 整体流程与双支路设计整个网络大致可以分成三个模块特征提取器、双流 MIL 聚合结构、融合分类头。特征提取器一般用 ImageNet 预训练的 VGG16 或者 ResNet把每一个 patch 编码成一个固定维度的特征向量。这里有个细节值得注意patch 之间的空间位置信息在基础版特征提取里是不保留的MIL 模型把整张 WSI 看成是无序的包这也是多实例学习“排列不变性”的天然要求。双流结构分上下两条支路。上支路走 embedding-based 路线把 bag 内所有实例的特征向量输入一个注意力聚合模块得到 bag-level embedding再接一个全连接层输出包的 logit。下支路走 instance-based 路线每个实例的特征先经过一个全连接层输出实例级 logit然后通过高斯混合模型GMM对这些 logit 的分布建模依据后验概率筛选出最具有判别力的 top-K 个实例把选出来的实例特征和 logit 组合成 bag-level 表示。两条支路不是完全独立的。其中一个关键设计是上支路的 bag embedding 和下支路挑出来的关键实例特征会拼在一起输入融合分类头得到最终的 slide-level 预测。这样上支路可以从全局角度捕捉整体组织模式下支路则强制模型关注那些最可疑的区域互相弥补。2.2 高斯混合模型与关键实例选择全篇最核心的一环论文里最吸引我的不是双流本身而是它如何用高斯混合模型来选择关键实例。这个设计背后的问题意识很明确在一张 WSI 的几万个 patch 里真正有判别力的往往只有少数几个大部分 patch 都是正常组织、背景、或者和肿瘤无关的结构。如果对全部实例一视同仁很容易被大量正常 patch 淹没如果只取 top-1 最大概率实例又太容易受到个别噪声样本的干扰。GMM 在这里就扮演了“动态证据筛选器”的角色。具体思路是这样的把每个实例在 instance 支路上输出的 logit 看成来自两个高斯分量的混合分布——一个分量对应阴性实例另一个对应阳性实例。通过期望最大化算法EM迭代估计两个高斯分量的均值、方差和混合系数。每次迭代里E 步根据当前参数计算每个实例属于阳性分量的后验概率M 步用这些后验概率重新估计参数。收敛后每个实例对应一个“阳性可能性”的软标签模型可以选择那些后验概率最高的实例作为关键实例进入后续的融合分类。我复现时的理解是GMM 在这里起到了两个作用一是提供了一种不需要额外监督的阈值自适应方法——与其固定取 top-K不如根据分布形状动态判断哪些实例真的偏离了阴性主分布二是给双流之间的信息传递提供了一个更平滑的加权方式——不是硬剪枝掉所有低分实例而是按照后验概率对实例特征做加权聚合保留了一部分“灰色地带”的信息。这里建议读原文时多花点时间看 EM 更新公式尤其是方差项的更新。实操中如果不加任何约束EM 很容易收敛到方差趋近于 0 的退化状态也就是某一个高斯分量只剩下一个实例这会让整个筛选机制失效。后面我会在训练避坑部分专门讲这个问题。2.3 为什么非用双流不可在我的实验直觉里只用 embedding-based 流模型从整体上判断“这张片子像不像癌”并不难难的是让模型说清楚到底哪一块区域触发了判断。普通 attention 聚合出来的权重分布往往比较平滑很多正常区域也分到了不低的注意力画出来的热图在临床上不能直接用来做病灶定位。只用 instance-based 流热图倒是有了比如每个 patch 都可以输出一个 0 到 1 的分数但包级精度通常在多个数据集上会比 embedding-based 低几个点。原因也很直白patch 级别做预测时局部特征区分度不够正常组织的某个腺体结构、某个炎症区域在低倍镜下看起来可能就是很像肿瘤当这些误判的 patch 数量多了包级聚合就会被带偏。双流设计的价值在于它让两个支路在同一个优化目标下相互矫正。instance 支路负责提供“哪里最可疑”的候选embedding 支路负责从全局角度对候选信息作出最终的、更稳健的判断。而且两路共享同一个特征提取器和大部分网络层并没有把模型体量翻一倍训练代价增加有限这种性价比在实际项目里非常重要。3. 实验设置与结果解读3.1 数据集与评估口径论文的验证场景选的是肿瘤检测任务我印象中最核心的数据集是 Camelyon16这个数据集包含几百张淋巴结转移癌的全切片图像slide 级别标签只区分“有转移”和“无转移”是弱监督 MIL 论文的标准基准之一。另一个常见的验证集是 TCGA 下的某个癌种数据。评估指标方面paper 的主角是 slide-level 的 AUC 和准确率这两种指标衡量的是“模型在整张 WSI 级别能不能正确判断是否有肿瘤”。此外由于 instance 支路天然输出每个 patch 的概率作者还会做 instance-level 的检测评估这里用的指标通常是 FROC 或平均精度。有个容易被新手忽略的点Camelyon16 的 WSI 里肿瘤只占整张片子很小的一部分可能不到总面积的 5%。这意味着在任何 patch 级别评估里绝大多数 patch 是阴性如果你用普通的准确率来衡量 patch 分类效果哪怕模型把所有 patch 都判成阴性准确率也能到 95% 以上毫无参考价值。所以一定要用带召回-精度折中关系的指标这一点在复现的时候要特别注意。3.2 关键结果与对比按照论文报告的结果Dual-stream MIL 在 Camelyon16 上的 slide-level AUC 大约在 0.93 左右具体数值以论文原文为准明显超过只用 max-pooling 的基础 MIL 方法和只用 attention pooling 的 embedding-based MIL 方法。更有参考价值的是和全监督模型的对比即使在 patch 级别完全没有肿瘤位置标注的情况下双流 MIL 的 slide 分类性能已经接近一些用了全监督分割标注的模型。这说明弱监督方法在这个任务上确实有落地价值。我把几种方案的定位整理成一个对照表方便后面做方案选型时参考方法监督级别slide-level 性能能否输出病灶热图训练稳定性全监督分割模型像素级高能依赖大量标注Max-pooling MILslide 级中等能但噪声大易受噪声样本干扰Attention MILslide 级较高间接、平滑相对稳定Dual-stream MILslide 级最高档能实例级概率需要调参后期稳定3.3 消融实验的作用论文里消融实验是理解设计动机的关键。我记得比较清楚的结论是两条去掉 instance 支路只保留 embedding 聚合slide-level AUC 会下降而且可视化热图明显变模糊去掉 GMM 的动态选择机制改成简单粗暴的 top-K 固定选择性能也会下降。这篇文章的消融实验最有价值的点在于证明了“GMM 动态筛选”不是锦上添花的加分项而是双流结构有效运转的核心。它像一个分配器告诉模型当前 bag 里哪些实例值得进融合阶段哪些实例应该被压制。没有这个机制双流中间的信息对接就变成了硬编码拼接失去自适应性。从我的角度看消融实验还间接说明了一个现象instance 支路和 embedding 支路之间存在一种互促进关系。EM 每轮根据当前模型输出重新估计分布筛选出更准确的阳性候选这些候选又通过融合层直接影响最终分类 loss反向传导时再优化整个特征提取器。整个过程像一个迭代优化的闭环比一次性的注意力加权更符合“逐步聚焦”的直觉。4. 复现与训练从论文到可以跑的代码4.1 数据预处理全套流程复现这篇论文最花时间的往往不是模型结构而是 WSI 预处理。我的大致流程是这样走的。第一步是用 OpenSlide 读取 WSI 并切 patch。要注意选择合适的切分倍率很多 MIL 论文用 20 倍物镜下的 patch尺寸选 256×256 或者 512×512。切完的 patch 要先做背景过滤最简单的办法是计算每个 patch 的 RGB 均值丢弃掉那些亮度极高、接近纯白区域的 patch因为那些是空白背景对训练没有任何贡献。第二步是染色归一化。不同医院、不同实验室出片的染色深浅差异非常大如果模型没做过染色归一化在新数据上性能会掉得很难看。论文里没有把染色归一化当成核心贡献但实际复现时几乎都会做。第三步是特征提取。把清洗后的 patch 输入预训练 CNN比如 VGG16 的 pool 5 输出得到固定维度的特征向量存成 h5 文件。这一步非常占磁盘空间一个好的思路是提前把所有图的 patch 特征提取好、缓存到本地训练过程中只读特征不再重复跑 CNN能大幅加快迭代。4.2 模型搭建核心代码示意以我复现时的理解Dual-stream MIL 的 forward 逻辑可以简化成下面这段 PyTorch 风格代码。这里不是官方源码的逐行复刻而是抓住核心流程的参考实现帮你快速跑通流程再逐步细调。import torch import torch.nn as nn from torch.distributions.normal import Normal class DualStreamMIL(nn.Module): def __init__(self, in_dim512, hid_dim128, num_select8): super().__init__() # instance 支路: 每个 patch 先映射到 logit self.instance_fc nn.Sequential( nn.Linear(in_dim, hid_dim), nn.ReLU(), nn.Linear(hid_dim, 1), ) # embedding 支路: 全局聚合 self.attention nn.Sequential( nn.Linear(in_dim, hid_dim), nn.Tanh(), nn.Linear(hid_dim, 1), ) # 融合头 self.classifier nn.Sequential( nn.Linear(in_dim num_select, 256), nn.ReLU(), nn.Dropout(0.5), nn.Linear(256, 1), ) def forward(self, x): # x: [B, N, D], N 为 bag 内实例数 B, N, D x.shape # instance 支路 inst_logits self.instance_fc(x).squeeze(-1) # [B, N] inst_probs torch.sigmoid(inst_logits) # 用 GMM 估计阳性/阴性分布(简化版: 以当前数据做单轮估计) # 实际更稳的做法是保存上一步的参数, 在训练 loop 里做多轮 EM mu_pos inst_probs.mean(dim1, keepdimTrue) mu_neg 1 - mu_pos var_pos inst_probs.var(dim1, keepdimTrue) 1e-6 var_neg var_pos # 简化假设两者同方差 pos_ll Normal(mu_pos, torch.sqrt(var_pos)).log_prob(inst_probs) neg_ll Normal(mu_neg, torch.sqrt(var_neg)).log_prob(inst_probs) posterior torch.exp(pos_ll - torch.logaddexp(pos_ll, neg_ll)) # [B, N] # 选出 top-K 关键实例 topk_idx torch.topk(posterior, knum_select, dim1).indices selected torch.gather(x, 1, topk_idx.unsqueeze(-1).expand(B, num_select, D)) selected_probs torch.gather(inst_probs, 1, topk_idx) # embedding 支路 attn_weights torch.softmax(self.attention(x).squeeze(-1), dim1) bag_emb (x * attn_weights.unsqueeze(-1)).sum(dim1) # [B, D] # 融合分类: 全局特征 关键实例特征 关键实例概率 combined torch.cat([bag_emb, selected.reshape(B, -1), selected_probs], dim1) logit self.classifier(combined).squeeze(-1) return logit, inst_probs, posterior这段代码是我在复现时为了跑通 ablation 而写的简版它抓住了双流加 GMM 筛选的关键流程但和论文原文的某些细节有出入。正式复现时强烈建议去读作者公开的源码尤其是 EM 每轮迭代的参数更新部分比我这里的单轮估计鲁棒得多。我这里写出来的意义是帮你理解计算图是怎么流动的不至于对着论文公式不知道从哪下手。4.3 训练参数与 trick训练阶段有几个经验值得记录。首先是学习率的设置backbone 特征提取器如果冻结只用较小的学习率比如 1e-4 到 3e-4 训练 MIL 头部如果选择微调 backbone学习率要降到 1e-5 以下否则很容易灾难性遗忘 ImageNet 学到的底层特征。第二点关于 bag 的大小。一张 WSI 切出来几万个 patch如果全塞进一个 bag 做训练显存压力非常大。常见的做法是采样一个子集比如每轮训练从全部 patch 中随机采样 256 或者 512 个组成当前 bag。子采样会带来一定的性能损失但可以让训练稳定很多。论文里 GMM 的选择机制也在一定程度上减轻了子采样噪声的影响——即使抽到一些阴性 patch 居多模型也会通过后验概率找回关键实例。第三点是两个支路的收敛速度不一致。instance 支路收敛得比 embedding 支路快前期容易出现 instance 支路已经过拟合到训练集的 patch 分布、而 embedding 支路还没拟合好的情况。我的做法是先单独用 instance 支路的 loss 做几个 epoch 的 warmup让 patch 预测有一定区分度以后再打开双流融合和 GMM 模块。这个技巧在多个 MIL 任务里都实测有效。第四点关于损失函数。如果只用一个包级交叉熵作为最终监督GMM 模块内部缺乏直接梯度信号EM 更新的好坏全靠分类 loss 间接反馈。我建议在 warmup 阶段额外加一个辅助的 instance 级损失比如强制 top-K 实例的预测概率接近包标签。这会显著加快收敛又能让 instance 支路不至于完全朝一个钝头的方向偏移。4.4 常见问题速查复现过程中最容易踩的坑我整理成了下面的表里面每一条都是我实际遇到过的现象可能原因解法显存 OOMbag 内实例数过多子采样到固定数量或者梯度累积EM 的方差变成 0高斯分量坍塌到单个实例对方差加下界约束或者对混合系数做拉普拉斯平滑训练初期 loss 不降GMM 初始化太差后验概率全偏向一类先用 instance 支路 warmup再开 EM 更新包级 AUC 高但热图很碎instance 支路过拟合 patch 噪声增大 dropout或用后验概率做阈值过滤后再出图不同染色风格下性能骤降没有做染色归一化预处理阶段加染色归一化例如 Macenko 方法正负 patch 极度不平衡肿瘤区域占比太小在子采样时做困难负样本挖掘保留部分高置信负样本5. 复现完之后的几点体会5.1 这个设计对后续工作的启发当我完整跑通这篇文章的复现流程后最大的感受是双流结构本身并不稀奇真正值得借鉴的是“通过一个无监督分布模型来动态筛选证据”的思路。后来的很多 MIL 工作包括基于注意力门控的 DSMIL、基于 token 筛选的 Transformer 类方法其实都在做同一件事——在弱监督条件下如何从海量 patch 里找到少数有效证据。Dual-stream MIL 用 GMM 找到了一个解释性很强的解决方案把阳性 patch 看成混合分布中的一小簇用统计模型去识别它们。我自己做过的实验里把同样的 GMM 思路迁移到其他弱监督分类任务上比如乳腺癌 HER2 状态预测也能观察到明显的稳定性和可解释性提升。这说明论文的核心贡献不完全绑定在 WSI 这个具体场景上而是一种更通用的“弱监督聚合 分布建模”思想。5.2 局限性与改进方向论文的局限性也很明显。GMM 的初始化对结果有一定敏感性如果两个高斯分量的初始均值设在同一点EM 更新可能很慢甚至陷入局部最优。作者用 GMM 建模 patch 后验分布本质上是假设 patch 概率服从高斯混合这个假设在真实组织切片里并不总成立特别是遇到大量炎症、坏死、出血区域时patch 概率可能呈现多峰分布两个高斯分量不够用。如果我来做改进会尝试三个方向一是用更灵活的分布模型替换 GMM比如基于归一化流或者潜变量模型来做证据筛选二是把 EM 更新过程改成可微的端到端模块让模型在训练过程中自适应地学习分布参数而非依赖外部迭代三是引入多尺度信息因为病理诊断本身依赖层级观察低倍镜看组织结构高倍镜看细胞形态单一尺度的 patch 在信息表达上天然受限。5.3 给后来者的实操建议最后分享一点个人的实操习惯复现论文时比起直接对着公式抠细节我建议先从代码把整体流程跑通再逐步把论文里的每个模块替换成自己的实现。先用一个小的 toy dataset 验证 data pipeline 没问题再用小规模 Patch 采样跑 20 个 epoch 看 loss 是否下降最后才在完整数据集上大规模训练。这样能避免你在第一次训练时就把时间浪费在“搞不清是模型 bug 还是数据问题”上。另外弱监督 MIL 模型的调试反馈周期比较长一个 bag 里的 patch 数量动辄几千训练一个 epoch 可能就要几分钟。建议所有中间结果都可视化出来每 5 个 epoch 输出一次当前模型的 top-K patch 热图贴在 WSI 原图上看看模型聚焦的区域是否符合直觉。这类“人工巡检”在弱监督任务里极其重要因为训练损失在下降不代表模型学到的是你想要的那个病灶完全有可能是在利用染色伪影或者其他背景线索做判断。
返回列表