☰
双流多实例学习:从全切片图像到弱监督肿瘤检测与定位
2026/10/5 12:21:32 网站建设 项目流程

拿到一张全切片图像(whole slide image,WSI),你手里唯一的监督信号是“这张片子有没有肿瘤”,连具体位置都不告诉你——这时该怎么训练一个能同时给出诊断结果和病灶热图的模型?这就是 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_dim=512, hid_dim=128, num_select=8): 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(dim=1, keepdim=True) mu_neg = 1 - mu_pos var_pos = inst_probs.var(dim=1, keepdim=True) + 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, k=num_select, dim=1).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), dim=1) bag_emb = (x * attn_weights.unsqueeze(-1)).sum(dim=1) # [B, D] # 融合分类: 全局特征 + 关键实例特征 + 关键实例概率 combined = torch.cat([bag_emb, selected.reshape(B, -1), selected_probs], dim=1) 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 原图上看看模型聚焦的区域是否符合直觉。这类“人工巡检”在弱监督任务里极其重要,因为训练损失在下降不代表模型学到的是你想要的那个病灶,完全有可能是在利用染色伪影或者其他背景线索做判断。

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

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

立即咨询