医学图像里的 CT 病灶检测,一直是“数据贵、标注难、模型不容易收敛”的典型场景。最近看到一篇很有意思的工作,标题里有几个关键词组合得非常巧妙:“Lesion Detection in CT with Frozen Self-Distilled Features: SALT, a Spatially Adaptive Label-Guided Temperature”。简单翻译过来就是:用冻结的自蒸馏特征做 CT 病灶检测,并且设计了一个叫 SALT 的空间自适应标签引导温度模块。
这篇文章我会按“为什么这么做 → 核心原理拆解 → 如何复现 / 如何做实验 → 常见坑 → 工程建议”的路线来讲。适合已经有 PyTorch 基础、想了解医学图像检测新方法的读者;如果你还没接触过自蒸馏,也不用担心,我会先补上概念,再进入 SALT 本身。
1. 病灶检测为什么需要“冻结自蒸馏特征”?
1.1 CT 病灶检测的难点在哪
CT(Computed Tomography,计算机断层扫描)影像在肺部结节、肝脏肿瘤、淋巴结筛查等任务中非常常用。但在深度学习落地时,有几个绕不开的问题:
- 标注成本极高:CT 是三维体数据,标注一个病灶往往需要医生逐层勾画边界,一个病例可能要花费几十分钟甚至更久。
- 类不均衡严重:病灶区域通常只占整个 CT 体积的很小一部分,背景像素占绝大多数。
- 数据分布差异大:不同品牌 CT 设备、不同扫描参数、不同重建算法都会带来灰度分布差异。
- 模型泛化难:在小规模数据集上训练的检测模型,换到新医院、新设备上性能下降非常明显。
所以,如何让模型在“有限标注”下学到更通用的特征,是这个领域非常关注的问题。
1.2 什么是自蒸馏特征,什么是“冻结”
在过去几年,自监督学习和自蒸馏是表示学习里的两个重要方向。这里我把它们放在一起解释。
自蒸馏(Self-Distillation),简单理解就是一个网络自己教自己。常见做法是把同一张图片做两次不同的随机增强,得到两个视角,然后让一个分支从另一个分支的特征中学习。经典的 DINO、EsViT、iBOT 等方法都使用了类似思路。
冻结(Frozen),指的是模型或者特征提取器的参数在后续任务训练中不再更新。比如我们先用自蒸馏在大规模数据上训练好一个 backbone,然后把它固定住,只训练后面的检测头或分割头。
为什么标题里特别强调“Frozen Self-Distilled Features”?因为这种做法有几个很现实的好处:
- 自蒸馏特征已经具备较强的语义和空间一致性,冻结后可以避免灾难性遗忘。
- 冻结 backbone 后,显存和计算量相对可控,可以集中资源训练检测头。
- 在不同下游任务间切换时,特征可以复用,适合多任务、多数据集场景。
1.3 SALT 是解决什么问题的
SALT 的全称是Spatially Adaptive Label-Guided Temperature,翻译过来是“空间自适应标签引导温度”。
只看名字其实比较抽象,拆开理解:
- Temperature(温度):在自蒸馏、知识蒸馏里,温度系数用来控制概率分布的平滑程度。温度越高,分布越平滑;温度越低,分布越尖锐。
- Label-Guided(标签引导):温度并不是全局一个固定值,而是根据标签信息来调制,让模型在不同区域用不同的“学习强度”。
- Spatially Adaptive(空间自适应):CT 是三维体数据,病灶在不同空间位置上的特征复杂度、标注置信度、样本难度都不一样,所以温度需要逐体素或逐区域地变化。
一句话总结:SALT 想解决的问题,是让冻结的自蒸馏特征在下游 CT 病灶检测任务中,通过一个空间变化的温度调度,更好地区分病灶和背景,尤其在小目标、边界模糊的情况下获得更稳健的表现。
2. 论文核心思路:SALT 的空间自适应标签引导温度
2.1 从“蒸馏温度”说起
如果你用过知识蒸馏,应该对下面这个软标签公式不陌生:
[ q_i = \frac{\exp(z_i / T)}{\sum_j \exp(z_j / T)} ]
其中 (z_i) 是 logits,(T) 是温度。当 (T=1) 时就是普通 softmax;当 (T>1) 时,输出分布更平滑,类间的“暗知识”会被放大。
在自蒸馏流程里,temperature 决定了 teacher 特征对 student 特征的影响程度。传统方法通常用全局常数或简单退火策略,但医学图像中病灶区域小、边界模糊,全局温度显然不够精细。
2.2 标签引导温度
“标签引导”是 SALT 的一个关键点。
在 CT 病灶检测任务中,我们通常有像素级或体素级标签,比如:
- 0:背景
- 1:病灶区域
如果所有位置的温度都一样,模型会均匀地对待每一个体素。但病灶区域和背景区域的“学习难度”完全不一样。于是 SALT 的思路是:
- 在标签为病灶的区域,希望模型更关注细节,温度可以更低,让预测更尖锐。
- 在标签为背景的区域,温度可以更高,让模型保持平滑,抑制假阳性。
当然,具体实现中不一定直接使用硬标签,也可能使用标签距离图、边界距离图、伪标签置信度等作为引导信号。这部分需要以论文原文或官方代码为准,但思想是清晰的:用标签信息去调制温度场,而不是让温度在全图统一。
2.3 空间自适应:为什么每个位置需要不同温度
病灶在 CT 中往往出现以下情况:
- 大小差异大,从几毫米到几厘米。
- 边界有的清晰,有的模糊。
- 有些与周围组织灰度接近,人眼都难以分辨。
- 三相/多期增强 CT 中,病灶在不同期相上的表现也不同。
如果用一个全局温度,边界模糊区域和背景区域容易混淆,造成漏检或误检。空间自适应就是要让温度成为一个“空间场”,每个位置都有自己的温度值:
[ T(x) = f(\text{feature}(x), \text{label}(x), \text{context}(x)) ]
这个 (T(x)) 可以是一个模块生成出来的,也可以由一个小的卷积网络预测出来。然后在损失函数中,不同空间位置使用各自的温度重新加权。
2.4 SALT 与整体训练流程的关系
从论文标题可以推断,训练流程大致可以拆成两阶段:
- 第一阶段:在大规模数据上通过自蒸馏学习通用特征,得到一个特征提取器。
- 第二阶段:冻结这个特征提取器,在 CT 病灶检测任务上训练检测头,同时在训练过程中引入 SALT 模块,用标签引导的空间自适应温度来调整损失权重和优化方向。
用文字描述就是:
输入 CT 体数据 -> 送入冻结的自蒸馏特征提取器 -> 得到体素特征 -> 检测头生成病灶预测 -> SALT 模块根据特征与标签生成空间温度场 -> 用温度场加权损失,反向传播更新检测头需要提醒大家:上面这个流程是我根据论文标题和通用自蒸馏框架做的合理推导。具体 SALT 模块的输入、输出通道数、损失函数形式,一定要以论文原文、官方代码或作者公开的实现为准。在没有确定源码之前,不要盲目把这里的伪流程当成论文精确实现。
3. 环境准备与实验资源建议
有了背景之后,我们来聊一聊如果想做类似实验,需要准备哪些环境和资源。
3.1 硬件与软件栈
CT 病灶检测通常处理三维数据,显存消耗比 2D 图像大很多。建议如下:
| 资源 | 建议 |
|---|---|
| GPU | NVIDIA RTX 3090 / A100 / V100,显存建议 24GB 以上 |
| CPU | 用于数据加载和预处理,建议多核 |
| 内存 | 至少 32GB,处理完整 CT 时需要同时缓存多个病例 |
| 存储 | CT 原始数据通常较大,需要 SSD 加速读取 |
软件栈方面,下面是一套很常见的组合:
- Python 3.8 或更高版本
- PyTorch 1.10 或更高版本
- MONAI:医学影像处理专用库,强烈推荐
- NumPy、SimpleITK、NiBabel:读取和预处理医学影像
- OpenCV / SciPy:辅助处理
- TensorBoard / wandb:实验记录
需要注意的是,版本号要结合你本机驱动和 CUDA 环境实际调整,不要直接照抄网上配置。如果 PyTorch 与 CUDA 版本不匹配,会出现CUDA error: no kernel image is available这类问题。
3.2 数据集与预处理
做 CT 检测,数据集一般来自医院内部、公开竞赛或合作单位。常见公开数据集有 LUNA16(肺结节)、DeepLesion(多类病灶)、KiTS(肾脏肿瘤)等,但公开数据的标注协议各不相同,实验时必须明确自己的评估指标。
CT 预处理通常包括:
- 重采样(Resampling):统一体素间距,比如统一为 1.0mm × 1.0mm × 1.0mm。
- 窗宽窗位(Windowing):根据病灶类型选取合适的窗位窗宽,比如肺窗、腹窗。
- 归一化:裁剪到指定 HU 范围后,线性缩放到 [0,1] 或 [-1,1]。
- 裁剪或分块:由于完整 CT 体积太大,通常会切成 patch 输入模型。
下面是一个基于 MONAI 的简易预处理片段:
import monai from monai.transforms import ( LoadImaged, EnsureChannelFirstd, Spacingd, ScaleIntensityRanged, CropForegroundd, RandSpatialCropd, Compose, ) train_transforms = Compose([ LoadImaged(keys=["image", "label"]), EnsureChannelFirstd(keys=["image", "label"]), Spacingd(keys=["image", "label"], pixdim=(1.0, 1.0, 1.0), mode=("bilinear", "nearest")), ScaleIntensityRanged( keys=["image"], a_min=-175, a_max=250, b_min=0.0, b_max=1.0, clip=True, ), CropForegroundd(keys=["image", "label"], source_key="image"), RandSpatialCropd(keys=["image", "label"], roi_size=(96, 96, 96), random_size=False), ])这里我把 CT 数值范围裁剪到-175到250,这是腹部和胸部比较常用的 HU 范围之一。你的数据集如果来自不同设备,这个范围要根据经验调整。
3.3 实验目录结构
建议用下面的目录结构组织实验,方便复现:
project/ ├── configs/ # 配置文件 │ └── salt_experiment.yaml ├── data/ │ ├── raw/ # 原始 DICOM/NIfTI │ ├── processed/ # 预处理后的数据 │ └── splits/ # 数据集划分文件 ├── models/ │ ├── backbone/ # 冻结特征提取器 │ └── heads/ # 检测头、SALT模块 ├── scripts/ │ ├── train.py │ ├── evaluate.py │ └── infer.py ├── runs/ # 日志和 checkpoint └── requirements.txt4. 用 MONAI + PyTorch 搭建一个可运行的特征冻结基线
在复现 SALT 之前,先搭建一个最简单的“冻结特征 + 病灶检测头”基线。如果你能跑通这个基线,再往里面加入空间自适应标签引导温度会容易很多。
下面我给出一个简化但可运行的框架代码,重点演示思路,不追求和 SALT 论文完全一致。
4.1 读取 CT 与标注
这里我使用带Label的 NIfTI 文件。如果没有现成数据,也可以用 MONAI 的合成数据来测试流程。
import os import numpy as np import torch from monai.data import DataLoader, Dataset from monai.transforms import Compose, LoadImaged, EnsureChannelFirstd, ScaleIntensityRanged data_dir = "data/processed" files = [] for case_id in os.listdir(data_dir): img_path = os.path.join(data_dir, case_id, "image.nii.gz") label_path = os.path.join(data_dir, case_id, "label.nii.gz") if os.path.exists(img_path) and os.path.exists(label_path): files.append({"image": img_path, "label": label_path}) transforms = Compose([ LoadImaged(keys=["image", "label"]), EnsureChannelFirstd(keys=["image", "label"]), ScaleIntensityRanged(keys=["image"], a_min=-175, a_max=250, b_min=0.0, b_max=1.0, clip=True), ]) dataset = Dataset(data=files, transform=transforms) dataloader = DataLoader(dataset, batch_size=1, shuffle=True, num_workers=4)注意,这段代码默认你的数据已经统一了 spacing,并且尺寸一致。实际项目中通常还需要Spacingd和RandSpatialCropd。
4.2 加载冻结特征提取器
特征提取器可以选择很多种:
- 在自然图像上自蒸馏的 ViT / Swin Transformer
- 在医学图像上预训练的 Swin UNETR encoder
- 自己用自监督/自蒸馏训练好的 3D CNN
关键点是:冻结参数。
import torch.nn as nn class FrozenBackbone(nn.Module): def __init__(self, backbone): super().__init__() self.backbone = backbone # 冻结所有参数 for param in self.backbone.parameters(): param.requires_grad = False def forward(self, x): # 冻结模式下,不计算梯度,节省显存 with torch.no_grad(): feats = self.backbone(x) return feats如果你的 backbone 里包含 BatchNorm,冻结后要小心一个问题:BatchNorm 在训练模式下会继续更新 running mean 和 running var。通常建议:
- 将 backbone 设置为
eval()模式。 - 或者把 BatchNorm 层也换成 FrozenBatchNorm。
下面是一个简单处理方式:
def freeze_batchnorm_stats(model): for module in model.modules(): if isinstance(module, torch.nn.BatchNorm3d) or isinstance(module, torch.nn.BatchNorm2d): module.eval()这样就能避免冻结 backbone 时 BatchNorm 统计量被下游任务数据带偏。
4.3 构建病灶检测头
检测头可以按你自己的任务选择:
- 如果是分割型检测:用
1x1x1卷积输出类别 logits。 - 如果是锚框检测:输出 box 回归和分类。
- 如果是点在点检测:可以用热图回归。
为保持示例简洁,我这里用体素分类的方式,也就是一个简易的 3D U-Net 风格检测头,输出每个体素是否为病灶的概率。
import torch.nn.functional as F class SimpleVoxelHead(nn.Module): def __init__(self, in_channels, hidden_channels=64, num_classes=1): super().__init__() self.conv1 = nn.Conv3d(in_channels, hidden_channels, kernel_size=3, padding=1) self.conv2 = nn.Conv3d(hidden_channels, hidden_channels, kernel_size=3, padding=1) self.out = nn.Conv3d(hidden_channels, num_classes, kernel_size=1) def forward(self, x): x = F.relu(self.conv1(x)) x = F.relu(self.conv2(x)) logits = self.out(x) return logits4.4 构建空间自适应标签引导温度简易示例
这里我给出一个简化版伪代码,用来演示“空间自适应温度”的思想。真实 SALT 的公式和结构需要以论文为准,我下面的写法是便于理解的示例:
class SpatialAdaptiveTemperature(nn.Module): """ 简易版本:根据标签和特征生成一个温度场。 这里使用 Sigmoid 将温度约束在合理范围。 """ def __init__(self, feature_channels, base_temperature=2.0): super().__init__() self.base_temperature = base_temperature self.temperature_head = nn.Sequential( nn.Conv3d(feature_channels + 1, 32, kernel_size=3, padding=1), # +1 为标签通道 nn.ReLU(inplace=True), nn.Conv3d(32, 1, kernel_size=3, padding=1), ) def forward(self, features, label): # 将 label 转为 float 并确保有通道维度 label = label.float() if label.dim() == 4: label = label.unsqueeze(1) inp = torch.cat([features, label], dim=1) # 输出温度增量,然后再叠加基础温度 delta = self.temperature_head(inp) temperature = self.base_temperature + torch.sigmoid(delta) * 4.0 return temperature这段代码思路是:
- 输入是当前体素特征和标签。
- 通过一个小卷积网络得到每个体素的温度增量。
- 用 Sigmoid 限制增量范围,避免温度过大或过小。
在真实 SALT 中,对温度的建模会更精巧,可能包括距离变换、多尺度特征、可学习温度上限等。这里的示例只是帮你建立直观理解。
4.5 训练循环伪代码
下面给出一个极简训练循环,演示如何把温度场加进损失函数。
from monai.losses import DiceLoss model = nn.Module() # 假设 model 已经包好了 backbone、head、temperature module optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4) loss_fn = DiceLoss(sigmoid=True) for epoch in range(50): for batch in dataloader: image = batch["image"].cuda() label = batch["label"].cuda() # 前向:这里需要按你的模型结构调整 logits, temperature = model(image, label) # 用温度对 logits 缩放,模拟“温度调制的概率分布” logits_modulated = logits / temperature loss = loss_fn(logits_modulated, label) optimizer.zero_grad() loss.backward() optimizer.step()这里有一个需要思考的点:我们是在损失函数中直接对 logits 除以温度,还是在生成 soft pseudolabel 时使用温度?不同设计会得到不同效果。SALT 论文中具体如何处理,需要通过原始实现确认。
5. 如果复现论文,需要在哪些环节落地 SALT?
论文系统性地做实验,通常会在下面几个环节重点设计。
5.1 Teacher 特征提取
冻结自蒸馏特征意味着我们要先准备一个“教师特征”或者“预训练特征提取器”。复现时你应该检查:
- 特征提取器是在什么数据上预训练的?2D 还是 3D?
- 自蒸馏的形式是什么?是否使用了多视角、多尺度?
- 特征维度是多少?能否直接对齐到检测头的输入通道?
如果特征提取器来自 2D 预训练模型,要处理 CT 三维体数据,通常有两种做法:
- 逐层切片提取 2D 特征,再组成 3D 特征。
- 用 2.5D 策略,多个正交平面分别提取特征后融合。
如果直接使用 3D 预训练模型,例如医学图像上的 Swin UNETR encoder,那么输入输出维度和感受野会更匹配。
5.2 损失函数与温度计算
这是 SALT 论文的核心。复现时要重点思考:
- 温度是加在哪个环节?是 loss 内部概率的缩放,还是特征的对齐权重?
- 标签引导是如何实现的?是硬标签、软标签、还是距离图?
- 空间自适应是逐体素、逐 patch、还是逐 slice?
- 温度场的生成模块是否和检测头一起端到端训练?
我建议先做一个最简版本:用DiceLoss + Focal Loss做基础损失,温度场作为调制项。然后逐步替换成论文方案,观察指标变化。
class CombinedLoss(nn.Module): def __init__(self): super().__init__() self.dice = DiceLoss(sigmoid=True) self.bce = nn.BCEWithLogitsLoss() def forward(self, logits, label, temperature=None): if temperature is not None: logits = logits / temperature return self.dice(logits, label) + 0.5 * self.bce(logits, label)这个混合损失在类不均衡的医学分割任务中比较常用。
5.3 评估指标
CT 病灶检测的评估指标通常包括:
| 指标 | 说明 |
|---|---|
| DICE | 病灶区域重合度 |
| Sensitivity / Recall | 查全率,关注漏检 |
| Precision | 查准率,关注误检 |
| F1-Score | 综合指标 |
| FROC / CPM | 自由响应 ROC,常用于结节检测 |
| 假阳性数 | 每例平均假阳性数量 |
如果阅读论文,建议重点关注作者在实验表里使用的指标。不同数据集和任务,指标选择会直接影响结论。
5.4 消融实验设计
复现 SALT 时,消融实验应该覆盖以下几个方面:
- 不加空间自适应温度,使用全局固定温度。
- 加空间自适应温度,但不使用标签引导。
- 使用标签引导温度,但不做空间自适应。
- 完整 SALT。
对比这些变体,才能确认每个组件的贡献。如果你自己也在做类似方法,建议把这一套消融流程固定下来。
6. 常见问题与排查思路
在实践这类方法时,很容易遇到下面这些问题,我整理了一份排查表:
| 问题现象 | 常见原因 | 解决思路 |
|---|---|---|
| 训练时显存不足 | 3D patch 过大或 batch size 过大 | 减小 patch 尺寸、减小 batch、使用梯度累积 |
| 冻结 backbone 后特征全为 0 | 输入归一化错误或 backbone 参数未加载 | 检查 CT 数值范围、打印特征统计量 |
| 训练不收敛 | 温度初始化不合理 | 尝试温度固定为 1.0 或 2.0,观察 loss |
| BatchNorm 统计量漂移 | 冻结层仍处于训练模式 | 对冻结层调用 eval() 或替换 FrozenBatchNorm |
| Dice Loss 为 NaN | 标签全为背景或梯度爆炸 | 加 smooth 项、检查标签是否为空、使用混合精度 |
| 验证性能低于预期 | 预处理与训练不一致 | 确保推理时使用同样的窗宽窗位和 spacing |
| 训练速度太慢 | 3D 数据读取和增强耗时 | 使用 MONAI CacheDataset、缓存预处理结果 |
6.1 一个具体案例:特征全为 0 的排查思路
如果你加载预训练模型后,发现提取出来的特征全为 0,建议按下面步骤排查:
# 第一步:打印输入统计 print(image.min().item(), image.max().item(), image.mean().item()) # 第二步:打印特征统计 with torch.no_grad(): feat = backbone(image) print(feat.min().item(), feat.max().item(), feat.mean().item())如果输入正常但特征全为 0,大概率是 pre-trained 权重没有正确加载,或者权重文件路径错误。不要急着调模型结构,先确认权重加载成功:
# 检查权重加载 ckpt = torch.load("pretrained_weights.pth", map_location="cpu") print(ckpt.keys())7. 工程化最佳实践
不管是在读 SALT 论文还是做自己的实验,下面的工程经验都值得保留。
7.1 关注数据安全与合规
医学影像数据涉及患者隐私,做实验时要有严格的合规意识:
- 不要将未脱敏的 DICOM 直接上传到公共仓库或网盘。
- 代码仓库中不要包含患者 ID、医院信息。
- 与医院合作时,确认数据使用授权和伦理审批。
- 涉及生产环境或临床辅助诊断时,必须获得相应法律和伦理许可。
在写博客、开源代码时,只使用公开数据集或合成数据。
7.2 结构化配置管理
实验参数多,建议用 YAML 管理:
# configs/salt_experiment.yaml data: data_dir: "data/processed" roi_size: [96, 96, 96] spacing: [1.0, 1.0, 1.0] model: backbone: "swin_unetr_encoder" freeze_backbone: true feature_dim: 48 head_hidden: 64 salt: base_temperature: 2.0 temperature_range: [0.5, 6.0] label_guided: true spatial_adaptive: true training: batch_size: 2 lr: 1e-4 epochs: 50 amp: true这样每次实验都能保留一组配置,很利于复现。
7.3 使用混合精度与分布式训练
3D 检测训练很慢,建议:
- 使用
torch.cuda.amp混合精度训练。 - 多卡训练时使用
DistributedDataParallel。 - 使用
MONAI的缓存机制减少数据加载瓶颈。
一个简单的 AMP 片段:
scaler = torch.cuda.amp.GradScaler() for batch in dataloader: image = batch["image"].cuda() label = batch["label"].cuda() with torch.cuda.amp.autocast(): logits = model(image) loss = loss_fn(logits, label) optimizer.zero_grad() scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()7.4 实验记录
在训练过程中记录以下几类信息:
- 损失曲线(Dice Loss、BCE Loss、总 Loss)
- DICE 和 F1 的变化
- 温度场的统计量:均值、方差、极小值、极大值
- 显存占用、每 epoch 耗时
建议每个实验都保留:
- 模型代码版本 / commit id
- 配置文件
- 随机种子
- 数据集划分文件
随机种子对医学图像实验影响比较大,建议固定:
import random import numpy as np import torch def set_seed(seed=42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed)8. 总结与下一步路线
读到这里,你应该对“CT 病灶检测 + 冻结自蒸馏特征 + SALT 空间自适应标签引导温度”这条技术路线有了整体认识。本文的价值在于帮你把论文标题中的几个关键概念拆解清楚,并给出一个可以落地的实验框架:冻结特征提取器、体素检测头、温度场调制、损失函数结合、常见问题排查。
如果你接下来想深入做这个方向,建议按这样的路线走:
- 先跑通 MONAI 的 3D 分割或检测基线,理解数据流和损失函数。
- 用公开 CT 数据集做预训练特征提取器的实验。
- 实现一个最简单的固定温度蒸馏基线,记录指标。
- 再实现空间自适应标签引导温度,逐步替换。
- 做消融实验,确认每个模块的贡献。
对于 SALT 论文本身,由于目前公开信息还比较有限,我的建议是:不要只依赖博客解读,一定要去读原始论文和官方代码。算法的具体公式、温度范围、标签编码方式、训练流程,只有原始实现才是最可靠的依据。如果你是做科研复现,可以围绕“温度场如何生成”“标签引导如何设计”“空间自适应如何实现”这三个问题去读代码,效率会高很多。
希望这篇文章能给你一些启发。如果你近期也在尝试类似的“冻结特征 + 医学图像检测”实验,欢迎收藏备用,按上面的步骤一步步搭起来。