论文标题:TransMIL: Transformer based Correlated Multiple Instance Learning for Whole Slide Image Classification
发表会议:NeurIPS 2021
源码地址:https://github.com/szc19990412/TransMIL
研究团队:清华大学、哈尔滨工业大学、北京大学
研究方向:计算病理学、全切片图像 (WSI)、多实例学习 (MIL)、Transformer、注意力机制
1. 研究背景与核心动机
在数字病理学中,全切片图像(WSI)的分类面临着图像尺寸巨大(千兆像素)且缺乏像素级标注的挑战。多实例学习(MIL)是目前解决此类弱监督学习问题的主流框架:将一个 WSI 视为一个包(Bag),将其切割成的小图像块视为实例(Instance),通过聚合实例特征来预测整个 WSI 的标签。
现有的 MIL 方法(包括基于注意力机制的 ABMIL、CLAM 等)通常建立在独立同分布 (i.i.d.) 假设的基础上,即认为包内的每一个图像块是相互独立的。然而,病理学家在进行诊断时,不仅会关注单一区域的局部形态,还会综合考虑不同区域之间的空间与形态学关联。因此,现有的 i.i.d. 假设并不完全符合临床病理的真实情况。
为了解决这一问题,研究团队提出了关联多实例学习(Correlated MIL)的理论框架,并基于此设计了 TransMIL,利用 Transformer 中的自注意力机制来探索实例之间的形态学与空间相关性。
2. 理论基础:关联多实例学习 (Correlated MIL)
论文在方法部分首先从数学理论上论证了 Correlated MIL 的可行性与优势。
定理 1 (Hausdorff 连续函数的近似):证明了一个连续的集合函数(打分函数
)可以被
形式的函数任意近似。这为 MIL 框架中引入特征映射和池化操作提供了理论依据。
推论 1:进一步扩展了定理 1,证明通过引入函数
(用于编码实例间的空间和相关性信息),打分函数同样可以被有效近似。
定理 2 (信息熵优势):这是最核心的理论贡献。证明了在“关联假设”下的实例信息熵
严格小于或等于“独立同分布 (i.i.d.) 假设”下的信息熵
。这在数学上严密地解释了为什么引入实例间相关性可以减少不确定性,从而为 MIL 问题提供更多有用信息。
基于上述理论,作者提出了一个通用的三步算法(Algorithm 1),并通过 Fig. 2 直观对比了不同池化矩阵
的区别:传统池化(Max/Mean)和基于旁路注意力的 MIL(如 ABMIL)对应的池化矩阵均是对角矩阵(忽略了实例间相关性);而基于自注意力的池化矩阵在非对角线位置存在非零值,显式地建模了实例间的相关性。
3. TransMIL 网络架构深度拆解 (Methodology)
为了实现 Correlated MIL,作者设计了 TransMIL 模型,其核心网络架构展示在 Fig. 3 (整体架构) 和 Fig. 4 (PPEG 模块) 中。
3.1 预处理与特征提取 (Input & Embedding)

-
输入:将 WSI 切割为
的图像块序列,丢弃背景区域,得到实例集合
,
为实例数量(每个 WSI 的
长度不一)。 -
特征提取:使用在 ImageNet 上预训练的 ResNet50 提取特征,再通过全连接层进行降维,最终将每个 WSI 转换为特征序列
(本文中
)。
3.2 TPT (Transformer and PPEG for TransMIL) 模块
TPT 模块是模型的主干,具体处理步骤对应论文中的 Algorithm 2:
步骤 1:序列重整 (Squaring)
-
目的:为了后续能够将一维序列还原为二维空间以提取局部位置信息,需要将长度为
的序列补齐为完全平方数。 -
操作:计算
,计算差值
。通过复制或其他方式补充
个特征向量。 -
引入 Class Token:在序列头部拼接一个用于聚合全局信息的 Class Token

-
输出:补齐后的新序列
。
步骤 2:第一次相关性建模 (Correlation Modelling)
-
操作:将
输入多头自注意力层 (MSA)。 -
Nyström 近似:由于 WSI 的实例数量
通常极大(数以万计),标准 Transformer 的复杂度为
,会导致内存溢出。TransMIL 采用了 Nyström 方法进行近似,选取
个地标节点 (Landmarks),将自注意力矩阵乘法的计算复杂度从
严格降低至
。 -
输出:经过形态学相关性建模后的特征
。
步骤 3:金字塔位置编码发生器 (PPEG, Fig. 4)

由于 WSI 序列长度可变,常规的绝对位置编码(Absolute PE)不适用。作者利用 2D 卷积自带的 Zero-padding 能够隐式编码位置信息的特性,设计了 PPEG(对应 Algorithm 3):
-
拆分:将
剥离出 Class Token,保留 Patch Tokens
。 -
2D 还原:将
重塑为 2D 图像特征图
。 -
分组卷积与融合:将
平行送入三个具有不同感受野的卷积核分支(
,并配合对应的 Zero-padding)。这一步不仅隐式编码了绝对位置信息,还通过不同粒度捕获了局部邻域的上下文信息。随后将三路输出与原输入
相加融合。 -
输出:展平回 1D 序列,再拼回 Class Token,得到融合了位置与局部信息的
。
步骤 4:深层特征聚合
-
操作:将携带了空间编码的
再次送入第二层 MSA。 -
输出:
。
步骤 5:映射与预测头 (MLP Head)
-
操作:提取输出序列的第 0 个位置特征(即经过充分信息交互的 Class Token),输入多层感知机 (MLP)。
-
输出:得到最终的包级别(WSI 级别)预测概率
。
4. 实验设计与结果分析
论文在三个公共数据集上进行了评估:CAMELYON16(二分类/不平衡)、TCGA-NSCLC(二分类/平衡)、TCGA-RCC(多分类/不平衡)。
4.1 WSI 分类性能 (Table 1)

-
结果对比:与传统池化方法及 ABMIL, DSMIL, CLAM 等最先进的 MIL 算法相比,TransMIL 在三个数据集上均取得了最佳的 Accuracy 和 AUC。
-
具体表现:在 CAMELYON16(肿瘤区域往往占比极小,
)上,TransMIL 的 AUC 达到 0.9309,比忽略实例相关性的模型高出至少 5%。在 TCGA-RCC 的多分类任务中,AUC 高达 0.9882,证明了模型对不平衡多分类任务的鲁棒性。
4.2 消融实验分析 (Table 2 & Fig. 5)

-
PPEG 模块有效性 (Table 2):作者评估了不同位置编码的效果。无位置编码 (w/o) 表现最差;正弦编码 (sin-cos) 和单一卷积核均能提升表现;而融合了多感受野的 PPEG 模块表现最佳。

-
条件位置编码验证 (Fig. 5):通过对比“输入序列保留原空间切割顺序 (order)”与“随机打乱输入序列 (w/o)”两组实验,证实了打乱顺序后 AUC 下降。这严密论证了 PPEG 能够有效捕获序列中蕴含的空间位置信息。
4.3 注意力可视化与可解释性 (Fig. 6)

-
模型输出的可解释性是临床应用的关键。Fig. 6 展示了 TCGA-RCC 数据集的切片。通过提取 TransMIL 第一层 Transformer 的自注意力权重并绘制热力图 (Attention heatmap),结果表明高注意力区域与病理学家精细标注(蓝色轮廓)的肿瘤区域高度吻合。
4.4 快速收敛特性 (Fig. 7)

-
相较于传统方法需要数十甚至上百个 Epoch 才能收敛,由于 TransMIL 同时显式地利用了实例间的形态学和空间相关性信息,模型学习效率大幅提高,其收敛速度和验证集 AUC 均显著优于 CLAM、DSMIL 等对比方法。
5. 总结
《TransMIL》这篇文章从数学理论出发,论证了 Correlated MIL 较于独立同分布假设在信息熵上的严格优势。在算法实现上,巧妙地利用 Nyström Transformer 解决了 WSI 长序列建模带来的 $O(n^2)$ 计算瓶颈,并通过定制的 PPEG 模块在变长序列上成功引入了条件位置与局部邻域编码。该框架在无需像素级标注的弱监督条件下,实现了 WSI 分类性能与可解释性的双重提升。

