ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

DSMIL双流多实例学习:弱监督下WSI肿瘤检测的关键技术解析

DSMIL双流多实例学习:弱监督下WSI肿瘤检测的关键技术解析 第一次翻到 Dual-stream multiple instance learning networks for tumor detection in Whole Slide Image 这篇论文时我正被一批淋巴结切片的肿瘤检测任务折腾得够呛。当时我已经在 Whole Slide ImageWSI上试了好几种常规分类网络但只要输入端变成一张几十亿像素的病理全切片所有标准化做法全都不管用了Resize 到 224×224 会丢掉肿瘤区域切成 patch 后一张片子轻松产生一两万个图像块标签却只有一个有无转移。这篇论文让我想明白了一件事——不能用处理自然图像的思路去硬啃 WSI而是要在 multiple instance learningMIL框架下用一套 dual-stream networks 把找关键区域和看全局上下文同时做掉。如果你也在做弱监督的病理图像分类或者被大图切块后标签怎么分配卡住这篇笔记应该能帮你少走不少弯路。这篇论文发在 CVPR 2020我记得作者是 Bo Li、Yin Li 他们那组。题目听起来不复杂但里面的双流结构设计、可微 Top-k 选择器、以及两路损失协同的训练方式拿到今天依然有不少值得复盘的地方。下面我按自己的理解拆开讲重点放在方法逻辑和复现时容易踩的坑上公式细节以原文为准这里只讲清楚为什么这么做。1. 为什么 WSI 上的肿瘤检测绕不开多实例学习1.1 WSI 的尺寸灾难一张图里藏着一座数据山先说说 WSI 这个输入到底有多离谱。一张标准的病理切片在 40 倍物镜下扫描图像分辨率通常超过 100000×100000 像素文件大小随压缩格式不同能到 2-4GB。普通自然图像分类用的 ResNet 输入是 224×224你把它直接喂进去就必须把整张图缩放到这个尺寸——这等于把一座城市压缩成一页地图。肿瘤灶可能只占整张切片的万分之一缩放之后直接消失模型看到的全是一片染色背景。更麻烦的是WSI 本身是多分辨率金字塔结构。底层是最高分辨率往上每层缩小两倍浏览器和查看器靠这种结构实现边拖动边加载。深度学习模型没法一路读完整个金字塔常规的做法是先在某一个倍率下把 WSI 切成小 patch比如 256×256 或 512×512然后对这些 patch 做分析。问题来了一张 40 倍率的淋巴结切片切完之后常常是几万个 patch其中真正包含肿瘤细胞的可能只有几十个到几百个正负样本极度不均衡。这种情况下如果做逐 patch 分类标签根本没法打。没有病理医生会给你逐 patch 标注这块是癌、那块是正常——一张切片人工标注区域级标签就得几十分钟面对几百上千张数据集根本不现实。能拿到的通常只有切片级别的报告有转移或者没有。这就决定了 WSI 分析天然要走弱监督路线。1.2 只有包的标签MIL 的天然适用场景多实例学习解决的就是这种只知道包标签、不知道实例标签的问题。在这个框架里一张 WSI 被看作一个 bag切出来的每个 patch 是 bag 里的 instance。训练时只给 bag 标签——这张切片有没有癌症模型要自己学会推断哪些 instance 在驱动这个标签。MIL 的标准假设是正包里至少有一个正实例负包里全部都是负实例。放在肿瘤检测场景里这个假设基本合理——病理报告说淋巴结有转移就意味着至少在某个视野里能找到癌细胞团。反过来如果报告说没有转移那这张切片里确实不存在明确的肿瘤灶。早期很多人直接用 max-pooling 做 MIL 聚合逻辑很直白把每个 patch 过一个分类网络得到一组成分数取最大值作为整张切片的分数。这个做法符合 MIL 假设但它有个明显毛病——它只关心响应最强的那个 patch多灶性肿瘤、分散的微转移灶全被忽略了。也有人用 mean-pooling把所有 patch 的响应平均一下但一张切片里绝大部分 patch 是正常组织均值一拉肿瘤信号被淹没在噪声里。DSMIL 这篇论文的思路就是在这个背景上做文章它既想保留找到最关键实例的能力又想兼顾全局上下文于是把两个流拆开来做。1.3 DSMIL 在 MIL 方法演进里的位置当时 MIL 在病理上的主流基线是 ABMILAttention-based MIL它用注意力机制给每个 patch 学一个权重再把特征加权求和。ABMIL 比 max-pooling 灵活但它学出来的权重是软的所有 patch 都会对最终表征有贡献正常组织的影响很难完全消除。后续的 CLAM 在 ABMIL 基础上加了实例级聚类约束TransMIL 则用 Transformer 编码全局关系——这些我后面会提到。DSMIL 的独特之处在于双流结构一个流用自注意力编码器关注实例之间的关系另一个流通过可微 Top-k 选出关键实例再做基于相似度的聚合。它的贡献不是发明了一个全新的聚合函数而是把实例判别和包级聚合两条路径解耦让模型同时具备两种能力。这在当时是一个很自然但没几个人做得干净的方向。2. DSMIL 的双流设计一个分支找关键一个分支看全局2.1 从整张切片到最终分数的完整流程DSMIL 的处理流程可以概括成四步预处理、特征提取、双流聚合、融合分类。预处理阶段在固定倍率论文用的好像是 20 倍下把 WSI 切成 patch过滤掉背景区域只保留有组织覆盖的 patch。特征提取阶段用一个预训练的 CNN比如 ResNet把每个 patch 编码成一个低维向量论文里通常还会对特征做 L2 归一化。之后这些 patch 特征进入双流模块第一流做自注意力编码得到实例级表征第二流通过可微 Top-k 找出关键实例并做包级聚合最后把两路的输出融合接一个分类头得到整张切片的预测分数。下面我拆开细讲。2.2 第一流自注意力编码器让 patch 之间互通消息第一流做的事情是把所有 patch 的特征序列送入一个自注意力编码器输出每个 patch 的上下文增强表征。为什么要自注意力因为病理组织中单个 patch 的形态往往具有歧义性。举个例子一个淋巴细胞团单独看可能觉得是反应性增生但如果它周围的 patch 全是结构破坏的异型细胞那它很可能是肿瘤浸润的一部分。自注意力让每个 patch 在计算表征时能看到其他 patch相当于给每个 patch 配备了全局视野。编码后的每个 patch 表征会接一个小分类头输出一个实例级分数。这一步非常关键——它让模型在训练过程中被隐式地推向识别出最可疑的 patch这个目标。不过单靠实例级分数直接做整图预测还不够鲁棒因为一个 patch 的响应可能有噪声所以需要第二流来兜底。2.3 第二流以关键实例为锚点的全局聚合第二流的输入同样是 patch 特征但它先通过一个可微 Top-k 模块选出 k 个最关键的实例。这里的关键由模型学出来的一个 score network 决定——每个 patch 会被打一个分排序后取前 k 个。选中这些关键实例后以它们的特征作为锚点计算全图所有 patch 特征与每个锚点之间的相似度再用 max-pooling 聚合。具体来说就是把与每个锚点最相似的 patch 特征挑出来和锚点特征拼接得到包级表示最后接分类头。这个设计的巧妙之处在于它不像 mean-pooling 那样把正常组织噪声平均进来也不像 max-pooling 那样只保留全局最大的一个响应。它选出来的锚点本身就代表模型认为最像肿瘤核心的区域然后通过全图相似度聚合把与这些区域模式相近的所有 patch 都捞进来看一眼。换句话说第一流负责找——定位可疑区域第二流负责联——把可疑区域与周边相关区域联合起来做最终判断。这种分工在思路上非常接近病理医生的读片流程先低倍率扫全片找到几个可疑区域再放大看这些区域以及它们和周围组织的关系。2.4 两流融合不是简单 ensemble我最初以为两流融合就是把两个分数平均一下读完论文发现不是这么回事。DSMIL 是把第一流生成的实例级特征和第二流生成的包级特征拼接或加权合并再过一个分类层做最终预测训练时两路损失一起反向传播。这意味着两流在训练过程中是互相影响的第一流的实例判别能力会影响 Top-k 选出来什么锚点锚点的质量又会影响第二流的聚合效果而第二流的全局信息也会反过来约束第一流不要只盯着孤立的高响应 patch。这种结构性互补比两个模型投票式的 ensemble 要高效得多。实际效果上单看第一流或单看第二流性能都会掉一截消融实验里能清楚看到这一点——后面我会专门讲。3. 可微 Top-k 选实例让挑重点也参与反向传播3.1 为什么需要 k 个而不是 1 个或全部稍微想过 MIL 的人都会问既然要找关键实例为什么不直接取分数最高的那一个原因有两个。第一肿瘤往往多灶性——一张切片里可能有多个独立的转移巢每个巢的位置、形态、细胞密度都不同只选一个最强的响应模型就只学会了找最大最典型的病灶微转移灶和小巢穴会被系统性忽略。第二最高分单个实例的响应可能包含噪声比如染色边缘伪影、组织折叠造成的假阳性取 k 个再做聚合能起到一定的鲁棒作用。但 k 也不能太大。k 接近 patch 总数时Top-k 退化为某种平均正常组织的噪声又回来了。论文里应该有对 k 的消融实验我自己的复现经验是 k 在一个中等范围比如 4 到 16内比较稳定再大后区分度明显下降。这个数本质上是在覆盖多灶性和过滤噪声之间找一个平衡点。3.2 硬 Top-k 的问题是梯度断流如果用 PyTorch 里的torch.topk来做这个选择得到的是索引索引操作本身不可导。也就是说score network 计算出的分数虽然是一个连续张量但一旦经过 topk 取索引、再用索引去 gather 特征梯度就没法从损失传回 score network 了——score network 成了死参数。这是很多人在自己做 MIL 时绕不开的坎所以早年的方法普遍用 soft attention 绕开硬选择。DSMIL 的思路是不回避硬选择而是把选择这个操作用可微的方式近似出来。这里有个关键认知我们不需要真的得到一个离散的 0/1 掩码只需要一个接近 0/1 但可导的掩码让梯度能够穿过选择过程。3.3 可微近似的核心逻辑这类方法一般通过两种思路实现。一种是基于排序松弛把 topk 操作视作对分数向量做排序再用连续函数比如 Sinkhorn 算子、NeuralSort 一类的 soft sort近似排序结果从而得到一个关于分数的连续掩码。另一种是 Gumbel-Softmax 式的采样松弛给分数加上 Gumbel 噪声再通过带温度参数的 softmax 得到近似 one-hot 的掩码。论文的具体推导我在笔记里标了回头再看复现时我用的是一个简化版score network 先给每个 patch 打分然后用带温度的 softmax 与可学习阈值做差得到一个软掩码让掩码和特征逐元素相乘后再 normalize。效果和论文描述的趋势一致——模型确实能学会把分数集中在少数关键 patch 上。这段经历让我明白一个道理可微 Top-k 的价值不仅在于能反向传播更在于它给了模型一种表达我就是要硬选几个重点的能力。软注意力再厉害本质还是对所有候选做加权求和它的焦点感和硬选择的聚合路径完全不同。3.4 k 值的选择经验k 值怎么定论文里应该画过 k 与性能的曲线。按我的经验k 太小训练不稳定尤其是初期 score network 还没学好的时候选出来的关键实例很可能是噪声k 太大双流结构逐渐退化到类似 mean-pooling两流的差异性就没了。如果项目数据和 CAMELYON 差异大建议从 k8 起步用验证集调另外训练初期可以把温度参数设大一点让掩码更平滑后期再减小温度让选择变得更硬这个 annealing 技巧能让训练稳很多。4. 损失函数组合与训练策略双流网络不撞车的关键4.1 两路损失怎么分配DSMIL 的损失大体由两部分组成第一流的实例级分类损失和第二流的包级分类损失两者加权相加。第一流的每个 patch 共享整张切片的 bag 标签——因为 MIL 假设正包里至少有一个正实例所以可以把 bag 标签复制给所有实例去监督训练当然绝大多数 patch 其实是负的这会带来不少噪声但配合第二流的包级约束模型能自己学会不把所有 patch 往正例上推。权重的设置直接影响训练行为。如果实例级损失权重过大模型会倾向于把大量 patch 都判成正类因为这样能快速降低 CE loss结果 Top-k 选出来的锚点全是滥竽充数的。如果包级损失权重过大第一流的实例判别能力变弱Top-k 可选不出来好东西。我复现时的做法是两路各 0.5 起步观察验证集 AUC 的变化再微调论文的实验部分也给出了他们最终用的设置但我建议把它当起点而不是终点不同数据集的最优配比差异挺大。4.2 特征提取器是整个流程的天花板这个坑我在别的 MIL 工作里也反复踩过很多人把精力全放在聚合器结构上忽略了特征提取器的质量结果模型上限被特征质量卡死。DSMIL 里 patch 特征来自预训练 CNN用什么预训练权重效果差别很大。ImageNet 预训练的 ResNet 能提供通用视觉特征但它没见过病理图像的纹理和染色模式特征里会有不少对病理无用甚至有害的维度。更优的选择是用自监督方法MoCo、DINO 这类在大量无标注病理 patch 上预训练的特征提取器或者用弱监督方式在相关任务上微调过一遍。这里给个量化的直觉同一个聚合器换一个更好的特征提取器AUC 可能从 0.85 涨到 0.90而你在聚合器上折腾一个月可能也就涨 0.01。所以做实验时第一件事就是固定一个优质特征提取器把变量控制在聚合方法上。特征维度方面512 或 1024 通常足够没必要上太大的向量反而增加注意力和 Top-k 的计算量。4.3 bag 采样、内存管理与训练节奏一张 WSI 生成的 patch 特征动辄上万条全部塞进双流模块显存扛不住。实际操作里是每个 batch 取若干个 bag每个 bag 随机采样固定数量的实例比如 2000 到 5000 个多余的部分丢弃或做重复采样。这里有个细节随机采样会改变实例的分布如果某张切片肿瘤区域特别小采样时可能一个肿瘤 patch 都没抽中导致这个 bag 变成假负样本。缓解办法是尽量在同一倍率下采样并且保证每个 bag 的采样数量足够大训练时多跑几个 epoch让不同 patch 的覆盖率达到。训练节奏上我建议先冻结特征提取器训练双流聚合部分等验证集 AUC 稳定后再解冻特征提取器做端到端微调学习率降到原来的十分之一。直接从头端到端训练特征提取器和聚合器会互相干扰损失曲线能看到明显的震荡收敛也非常慢。5. 实验验证维度AUC 之外还要看什么5.1 指标选择AUC、FROC 和它们的临床含义DSMIL 的实验主要在 CAMELYON16 和 CAMELYON17 上做这两个数据集都是乳腺癌淋巴结转移检测。CAMELYON16 的评估指标常用 AUC但真正贴近临床的是 FROCFree-response ROC曲线——它统计的是在每张切片假阳性个数限定下的检出灵敏度。我做过一个对比表格方便理解指标回答的问题特点适合场景AUC正样本分数是否普遍高于负样本只看排序不看具体阈值模型选型、快速对比FROC在可接受的假阳性数量内能检出多少真病灶更接近临床使用方式辅助诊断系统的性能评估只看 AUC 很容易被曲线下面积高骗过去。有些模型 AUC 很高但它的高灵敏区对应的是每张切片 10 个假阳性——这种模型在真实辅助诊断里根本用不了因为病理医生会被假阳性提醒淹没。所以复现时不能只盯着论文报的 AUC要自己画 FROC 曲线看低假阳性区间内模型到底什么表现。5.2 消融实验到底消掉了什么DSMIL 论文里最值得读的部分是消融实验。我当时总结出三个问题去掉第一流、把可微 Top-k 换成硬 Top-k 或 mean-pooling、去掉相似度聚合只保留锚点特征各自的性能变化是什么。这些实验的直接价值不是证明 DSMIL 全面领先而是告诉我们每个设计要素各自承担了什么职责。按我的理解消融结果大致呈现出这个规律没有第一流包级聚合缺了实例级判别的引导Top-k 选锚点的质量下降没有可微 Top-k换成硬选择训练收敛变慢甚至不收敛因为梯度传不到选择器把相似度聚合换成简单的锚点特征拼接模型对多灶性肿瘤的敏感度下降。这说明双流和 Top-k 都不是锦上添花而是互相咬合的整体设计。5.3 泛化性换个扫描仪还能不能打病理切片有一个自然图像领域不常遇到的麻烦不同医院的切片染色方案、扫描仪型号、扫描倍率都不一样导致同一个组织在不同数据集上颜色和清晰度差异很大。DSMIL 在 CAMELYON16 上表现好不代表在另一批用不同扫描仪采集的切片上也能直接复现那个 AUC。读论文时要注意它有没有做跨数据集验证或者用了什么染色归一化/数据增强手段。我自己的经验是训练时加入一些颜色抖动hue、saturation、brightness 扰动或者染色归一化预处理能显著提升换数据源后的稳定性。DSMIL 本身是一种结构设计它对染色差异并没有天然的免疫能力所以把它用到自己的数据上之前先做一次小规模跨域测试比什么都重要。6. 读完这篇论文我实际复现时踩过的坑6.1 patch 提取阶段背景过滤不是小事复现 DSMIL 的第一步就是切 patch但切这件事比想象中讲究。背景过滤阈值设高了会把边缘区域、染色很浅的组织当背景丢调设低了大量空白 patch 进入训练白白增加计算量。我自己一开始用的阈值偏高结果某些切片上肿瘤区域如果染色非常浅整个区域被当成背景滤掉标签还是正例模型自然学不好。建议先可视化几张过滤后的 mask看看组织覆盖率和保留区域是否合理。另外切 patch 的倍率也要先定好20 倍下 patch 数量可控特征语义更宏观40 倍下分辨率更高但一张切片会产生数万到十几万个 patch训练时间成倍增加。如果做的是检测任务且算力有限从 20 倍开始通常是最稳的选择。6.2 双流训练不稳定loss 权重和 top-k 的联动调整我第一次训练 DSMIL 时验证集 AUC 在前期一直不涨后来定位到两个问题一是两个流的 loss 权重失衡实例级损失主导后模型把所有 patch 都往正类推二是 score network 训练不充分Top-k 选出来的锚点在早期基本是随机的第二流学到的是噪声模式。解决办法不复杂先把 Top-k 的 k 设大一点让锚点集合更冗余给 score network 更多学习空间等训练中期分数分布稳定后再调小 k。同时给 score network 单独设一个偏小的学习率或用 warmup 让它先收敛到合理区域。这几个小调整做完收敛速度肉眼可见地变快。6.3 与已有 MIL 代码库的整合差异现在公开的 MIL 病理代码库很多最常被拿来当 baseline 的是 CLAM。代码写得清楚但它的核心聚合方式是 attention pooling直接在上面改 DSMIL 需要注意两点第一CLAM 的 bag 采样逻辑对 instance 数量有上限要求而 DSMIL 的第二流要做全图与锚点的相似度计算内存占用随 patch 数量线性增长不能照搬它的 batch 策略第二CLAM 的实例级聚类损失和 DSMIL 的实例级 CE loss 目标不同不能简单加在一起。我的建议是直接在 DSMIL 官方代码上改而不是从 CLAM 迁移否则很多潜在 bug 会消耗大量时间。6.4 双流结构能迁移到什么场景虽然这篇论文的标题是肿瘤检测但关键实例发现 全局上下文建模这个思路不限于 WSI。任何输入是大规模弱监督集合、只有集合级标签、且内部存在少数关键元素的任务都可以考虑这个框架。比如胶囊内镜图像中找异常区域、遥感影像中检测小目标、甚至文本分类里从一篇文章中挑关键句。我后来在一个组织切片的生存分析任务上试过类似结构把 bag 级分类换成 Cox 损失整体训练逻辑依然成立。路径上还可以做不少扩展把第一流的自注意力编码器换成更强的 Transformer 变体、把第二流的相似度度量改成学习式距离函数、或者在 Top-k 选择器中加入空间位置信息。DSMIL 最大的价值在于提供了一个稳定的骨架剩下的优化空间其实很大。我自己的体会是读 MIL 类论文时最忌讳只盯住最后的数字真正值得花时间的是弄明白每个模块为什么需要存在。DSMIL 的双流设计之所以能打动我是因为它把病理医生读片时的天然策略先扫视全局再聚焦可疑区域翻译成了可微分的网络结构。这种从真实流程里生长出来的设计比堆叠一堆新模块要扎实得多。如果你也想在自己的弱监督任务里用上这套思路建议先把论文的消融实验完整复现一遍再去改自己的数据——这个过程中你会对 Top-k 的敏感性、两路损失的平衡、以及特征质量的天花板有切身的体会。
返回列表