MICCAI 2021 的 GAMMA 挑战赛(Glaucoma Analysis on Multi-Modality imAges)任务一,官方给了两个东西:一份多模态眼底影像数据集,和一套基准代码(Baseline)。数据集里同一只眼睛既有彩色眼底照(CFP),也有视盘区扫出来的 OCT,标签是 0 到 4 的五级青光眼严重程度。Baseline 代码本身不复杂,几十行模型定义加一个常规训练循环,但它把"多模态数据怎么组织、两路图像怎么送进网络、融合层放在哪、指标怎么算"这条链路完整跑通了。很多人拿到官方仓库之后直接开跑,结果发现 loss 不降、验证集 AUC 在 0.5 附近打转,或者报一个形状不匹配的错就卡住。这篇内容就是把这套 Baseline 从头到尾拆开讲一遍:每个模块为什么这么写、参数为什么这么设、哪里最容易出问题,以及我实际跑过之后觉得值得改的地方。适合刚接触医学影像分类、准备复现或改进这套 Baseline 的同学,也适合想搞明白多模态融合代码落地长什么样的开发者。
1. 赛题拆解:多模态青光眼分级到底在考什么
1.1 任务定义与标签体系
GAMMA 任务一的输入是一组配对的眼底影像,输出是一个 0 到 4 的整数标签。这个标签描述的是青光眼的进展程度,等级越高代表视神经损伤越严重。它本质上是一个有序分类问题(Ordinal Classification),而不是普通的互斥分类,因为等级之间存在明确的顺序关系。这一点非常关键,后面选损失函数和评估指标的时候会反复用到。
官方数据按中心划分,训练、验证、测试的划分是固定的,不允许自己重新随机洗牌。这条规则不是形式主义,而是因为这个数据集的多中心特性非常强——不同中心的采集设备、成像参数、色温、分辨率都不一样,随机划分会让同一个中心的图像同时出现在训练和验证里,指标虚高得离谱,实际提交上去直接掉十几个点。我第一次跑的时候偷懒把训练和验证合并重新分了,本地 macro AUC 到 0.9 以上,换成官方划分后立马回落到 0.7 出头,这就是域偏移给上的第一课。
标签分布上,中间等级样本相对多一些,两端(0 级和 4 级)样本偏少,属于典型的长尾。Baseline 里没有做重采样,所以如果你直接跑,会发现模型倾向于预测中间类,混淆矩阵对角线两端全是空的。这是后面要重点处理的地方。
1.2 为什么是多模态:CFP 与 OCT 各自的价值
单看彩色眼底照,能观察到视盘形态、杯盘比、神经纤维层缺损这些结构线索,信息量大但受成像质量和医生主观判断影响明显。OCT 提供的是断层结构,能定量反映视网膜神经纤维层厚度、视杯深度这类参数,客观性强,但视野范围窄,只覆盖视盘及其周边一小块。
这两者不是冗余关系,而是互补关系。CFP 给"面"上的全局形态,OCT 给"深度"上的定量结构。青光眼早期,CFP 上可能看不出明显异常,但 OCT 的厚度图已经开始变薄;到了中晚期,CFP 上的杯盘比变化又比单张 OCT 更直观。所以题目逼着参赛者做多模态融合,本质是想看你能不能把两类证据整合起来,而不是简单地用一路图像刷分。
理解了这一点,融合层的设计就有了方向:它需要让两路特征在语义层面互相校正,而不是把两个向量首尾一接就完事。Baseline 用的是最朴素的拼接,属于"能跑就行"的版本,这也是它留给参赛者的改进空间。
1.3 官方 Baseline 的定位与整体设计取舍
官方 Baseline 的目标很明确:提供一个最小可运行的多模态分类范例,把数据读取、双分支前向、损失计算、指标评估这条链路完整展示出来,不追求分数。所以它做了几个明显的简化取舍。
一是 OCT 只取一张代表性 B-scan,或者把整个 volume 简单平均成单帧,而不是用 3D 卷积或序列模型处理整个 volume。二是融合方式用最简单的通道拼接加全连接,没有做注意力、没有做门控。三是数据增强几乎只有随机翻转和缩放,没有针对眼底图像做特定处理。四是损失函数就是普通交叉熵,没有处理类别不平衡。
这几个取舍在工程上完全合理——Baseline 的价值在于"可复现、易修改",不在于"高分"。你把这四个点里任意一个改好,都能拿到明显的涨分。所以我后面讲代码的时候,会同时告诉你官方是怎么写的、以及我改成了什么。
2. 数据管线:多模态眼底影像的读取与预处理
2.1 数据目录结构与标签对齐
官方仓库里数据一般按下面这种结构组织,训练集和验证集各自有独立的文件夹,标签放在一个表格文件里,常见是 Excel 或 CSV 格式。这里最容易踩的坑是索引对不上:图像文件名、表格里的行、以及最终 DataLoader 返回的样本顺序,三者必须严格一致。
GAMMA/ ├── train/ │ ├── CFP/ # 彩色眼底照,jpg 或 png │ ├── OCT/ # OCT 图像,可能是多帧命名 │ └── label.xlsx # 或 csv,含图像 ID 与分级标签 ├── valid/ │ ├── CFP/ │ ├── OCT/ │ └── label.xlsx └── test/ ├── CFP/ └── OCT/标签表格里通常有一列是样本 ID,一列是分级。OCT 的命名比较特殊,同一个病例可能有多张 B-scan,命名上会带序号。如果你的代码里用os.listdir直接读文件列表,然后用下标去索引标签,早晚会出事——os.listdir的返回顺序是文件系统决定的,不保证和后缀排序一致。
我习惯的做法是先把标签表格读进来,以它为准构建样本列表,然后去检查每一条对应的 CFP 和 OCT 文件是否真实存在,把缺失的样本直接剔除并打印出来。这一步花不了两分钟,但能省掉后面几个小时的困惑。实测官方数据里确实有个别样本的某一路图像缺失,如果不在读取阶段处理,跑到一半才崩,定位成本很高。
注意:不要用文件名做排序后直接当索引,一定要以标签表格为主表,做一次左连接式的存在性校验。
2.2 CFP 与 OCT 两路图像的预处理策略
两路图像的物理特性完全不同,预处理不能一刀切。
CFP 是彩色图,三个通道,视野是圆形或矩形,边缘有黑色背景。标准做法是:先统一尺寸,但不要直接拉伸,而是保持纵横比缩放后做中心裁剪或填充。眼底图像的视盘位置在不同图像里有偏移,直接拉伸会把圆形的视盘压成椭圆,杯盘比这种几何特征就被破坏了。我一般缩放到 512×512 并用零填充补齐,或者直接中心裁剪到有效视野区域。
另外,眼底图像经常偏色偏暗,特别是不同中心之间色温差异明显。业内常用的做法是 CLAHE 做局部对比度增强,或者做一次 gamma 校正把暗部细节拉出来。这里的 gamma 校正不是图像生成里那种噪声调度,而是纯粹的亮度映射,公式是 $I_{out} = 255 \cdot (I_{in}/255)^{\gamma}$,$\gamma$ 取 0.8 到 1.2 之间做微调。我在验证集上试过,加了 gamma 校正再配 CLAHE,早期青光眼的召回率能涨一点,但要注意别过度增强,否则噪声也被放大。
OCT 是灰度图,单通道。这里的关键问题是:一个病例对应多张 B-scan,怎么变成网络能吃的输入。三种常见方案,我列个表对比一下。
| 方案 | 做法 | 优点 | 缺点 |
|---|---|---|---|
| 单帧选取 | 取视盘中心附近最有代表性的一帧 | 计算量最小,实现简单 | 信息损失大,受选帧策略影响明显 |
| 多帧平均 | 把整个 volume 逐像素平均 | 降噪效果好,输入固定 | 抹平了层间结构差异 |
| 多帧堆叠 | 取固定帧数堆成通道或序列 | 信息保留完整 | 需要网络结构配合 |
官方 Baseline 大概率用的是前两种之一。我个人的做法是取中心连续 5 到 7 帧,缩放到统一尺寸后在通道维堆叠,然后用一个 3D 卷积或者逐帧 2D 卷积加时序池化的方式处理。如果只是想先跑通 Baseline,用多帧平均就够了,把单通道复制成三通道,直接塞进后面的双分支网络。
这里有个细节:OCT 和 CFP 的尺寸通常不一致,Baseline 里会对两路分别做 resize,最后输入网络的张量形状是(B, 3, H, W)各一路。你要确保两个分支的输入尺寸和主干网络的期望一致,ResNet 系对 224 友好,EfficientNet 系也是,别为了保持细节强行上 1024,显存顶不住。
2.3 Dataset 与 DataLoader 的实现要点
把上面这些整合成 Dataset 类,核心就三件事:读两路图、做增强、返回字典。下面这段是我改过的版本,保留了官方的骨架但补了几个关键处理。
import os import cv2 import numpy as np import pandas as pd import torch from torch.utils.data import Dataset, DataLoader from torchvision import transforms class GammaMultiModalDataset(Dataset): def __init__(self, df, root, cfg, is_train=True): self.df = df.reset_index(drop=True) self.root = root self.cfg = cfg self.is_train = is_train # 以标签表为准,过滤掉任一模态缺失的样本 cfp_dir = os.path.join(root, "CFP") oct_dir = os.path.join(root, "OCT") valid_idx = [] for i, row in self.df.iterrows(): sid = str(row["id"]) c_ok = os.path.exists(os.path.join(cfp_dir, f"{sid}.jpg")) o_ok = os.path.exists(os.path.join(oct_dir, f"{sid}.jpg")) if c_ok and o_ok: valid_idx.append(i) else: print(f"[warn] missing modality for {sid}, dropped") self.df = self.df.iloc[valid_idx].reset_index(drop=True) self.cfp_norm = transforms.Normalize( mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) self.oct_norm = transforms.Normalize( mean=[0.5], std=[0.5]) def _clahe(self, img): clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8)) return clahe.apply(img) def _gamma(self, img, gamma=0.9): inv = 1.0 / gamma table = np.array([((i / 255.0) ** inv) * 255 for i in range(256)]).astype("uint8") return cv2.LUT(img, table) def __len__(self): return len(self.df) def __getitem__(self, idx): row = self.df.iloc[idx] sid = str(row["id"]) label = int(row["label"]) cfp = cv2.imread(os.path.join(self.root, "CFP", f"{sid}.jpg")) cfp = cv2.cvtColor(cfp, cv2.COLOR_BGR2RGB) cfp = cv2.resize(cfp, (self.cfg["cfp_size"], self.cfg["cfp_size"])) oct_img = cv2.imread( os.path.join(self.root, "OCT", f"{sid}.jpg"), cv2.IMREAD_GRAYSCALE) oct_img = cv2.resize(oct_img, (self.cfg["oct_size"], self.cfg["oct_size"])) oct_img = self._gamma(oct_img, gamma=0.9) oct_img = self._clahe(oct_img) oct_img = np.stack([oct_img] * 3, axis=-1) if self.is_train: if np.random.rand() < 0.5: cfp = cfp[:, ::-1].copy() oct_img = oct_img[:, ::-1].copy() cfp = torch.from_numpy(cfp).permute(2, 0, 1).float() / 255.0 oct_img = torch.from_numpy(oct_img).permute(2, 0, 1).float() / 255.0 cfp = self.cfp_norm(cfp) oct_img = self.oct_norm(oct_img) return {"cfp": cfp, "oct": oct_img, "label": torch.tensor(label)}几个容易忽略的点说一下。第一,返回字典而不是元组,调试的时候能按名字取值,比记住下标顺序靠谱得多。第二,np.random.rand()做增强时两路必须用同一个判断结果,否则左右翻转后 CFP 和 OCT 的空间对应关系就错乱了——多模态任务里两路图像的空间对齐极其重要,这一点后面还会提。第三,OCT 复制成三通道是因为主干网络预训练权重是 ImageNet 的,通道数得对上;你也可以改网络第一层的卷积让它接受单通道,但那样预训练权重的第一层就得丢掉,小样本下不划算。
DataLoader 这边,num_workers设成 4 到 8,pin_memory=True,训练时shuffle=True,验证时shuffle=False。如果你用的是官方固定划分,验证集的顺序也是固定的,方便复现。
3. 模型结构:双分支编码器与融合方式
3.1 主干网络选型
Baseline 里两路通常共享同一个主干结构,只是分别初始化。选型上要考虑三点:数据量、显存、预训练权重可得性。GAMMA 的任务一训练样本只有几百例,这个量级禁不起大模型折腾。
我在实操里试过 ResNet-18、ResNet-50、EfficientNet-B0 三种。ResNet-50 参数量大,小样本上过拟合明显,验证集波动很大;ResNet-18 稳定但天花板不高;EfficientNet-B0 参数量约 5M,配上预训练权重效果最好,训练也快。最后我用的方案是两路各接一个 EfficientNet-B0,CFP 分支输入三通道,OCT 分支输入三通道(灰度复制),两路权重独立不共享。
为什么不共享权重?因为两路模态的统计分布差异太大,共享编码器会强迫它们映射到同一特征空间,反而拖累收敛。这个问题我做过对照实验:共享权重的版本验证 macro AUC 比独立权重低 3 个点左右,训练时长还多了将近一倍,因为共享权重下梯度要同时兼顾两路,收敛更慢。
加载预训练权重的时候记得处理分类头不匹配的问题,一般的写法是这样:
import torch.nn as nn from torchvision import models def build_backbone(name="efficientnet_b0", pretrained=True): if name == "efficientnet_b0": weights = models.EfficientNet_B0_Weights.IMAGENET1K_V1 if pretrained else None net = models.efficientnet_b0(weights=weights) feat_dim = net.classifier[1].in_features net.classifier = nn.Identity() # 去掉原分类头,输出特征 elif name == "resnet18": weights = models.ResNet18_Weights.IMAGENET1K_V1 if pretrained else None net = models.resnet18(weights=weights) feat_dim = net.fc.in_features net.fc = nn.Identity() else: raise ValueError(f"unsupported backbone: {name}") return net, feat_dim这里把分类头换成Identity,让主干直接吐特征向量,后面统一接融合模块。这么做的好处是主干和融合层解耦,你换主干的时候融合层代码一行都不用动。
3.2 融合层:从简单拼接到注意力
这是整个 Baseline 里最值得动刀的地方。官方版本基本就是拼接,可以写成这样:
class ConcatFusion(nn.Module): def __init__(self, feat_dim, num_classes=5, dropout=0.3): super().__init__() self.head = nn.Sequential( nn.Linear(feat_dim * 2, 256), nn.BatchNorm1d(256), nn.ReLU(inplace=True), nn.Dropout(dropout), nn.Linear(256, num_classes), ) def forward(self, f_cfp, f_oct): return self.head(torch.cat([f_cfp, f_oct], dim=1))拼接的假设是两路特征同等重要、彼此独立。但实际情况是,OCT 在中晚期更关键,CFP 在早期和整体形态上更关键,两者权重应该随样本变化。所以我改成了门控式的注意力融合:先算一个样本级的权重向量,再用它加权求和。
class GatedFusion(nn.Module): def __init__(self, feat_dim, num_classes=5, dropout=0.3): super().__init__() self.gate = nn.Sequential( nn.Linear(feat_dim * 2, feat_dim), nn.Sigmoid() ) self.head = nn.Sequential( nn.Linear(feat_dim * 2, 256), nn.BatchNorm1d(256), nn.ReLU(inplace=True), nn.Dropout(dropout), nn.Linear(256, num_classes), ) def forward(self, f_cfp, f_oct): g = self.gate(torch.cat([f_cfp, f_oct], dim=1)) # (B, D) f_cfp_w = f_cfp * g f_oct_w = f_oct * (1.0 - g) return self.head(torch.cat([f_cfp_w, f_oct_w], dim=1))这个门控的直觉是:gate 输出接近 1 时模型更信任 CFP,接近 0 时更信任 OCT。你可以把训练好的 gate 输出统计出来看分布,我在验证集上观察过,早期样本的 gate 均值确实偏大(更依赖 CFP),晚期样本偏小,和临床认知是吻合的。这种可解释性在医学影像任务里挺有价值,写报告的时候能拿来说事。
如果你想再进一步,可以做跨模态注意力:把 CFP 特征当 Query,OCT 特征当 Key 和 Value,算一次注意力。但要注意,小样本下注意力模块很容易训不动,我建议先用门控,等 baseline 稳了再上注意力。
3.3 分类头与损失函数设计
分类头就是两层全连接加一个 dropout,输出 5 维 logits。损失函数是重点。
普通交叉熵在这里有两个问题。一是类别不平衡,中间类样本多,两端少,模型会偷懒。二是忽略了标签的有序性,把"真实是 4 级预测成 0 级"和"真实是 4 级预测成 3 级"同等对待,这在临床上显然不合理。
我试过三种处理方式。第一种是给交叉熵加类别权重,权重取频率的倒数再归一化,实现简单,能缓解部分不平衡。第二种是 Focal Loss,对易分样本降权,让模型关注难样本,在小样本长尾场景下效果不错。第三种是软标签加 KL 散度,把硬标签按等级距离摊成分布,比如真实 4 级,标签可以写成 [0, 0, 0.1, 0.3, 0.6],这样模型预测成 3 级受到的惩罚比预测成 0 级小。
我最后采用的是第二种加第三种结合:主干用 Focal Loss 保证难样本被关注,同时把标签做一次邻域平滑,缓解有序性问题。参数上 Focal Loss 的 $\gamma$ 取 2,类别权重按频率开根号后取倒数,比直接取倒数更温和一点,避免极端权重导致训练震荡。
| 损失函数 | 适用场景 | 我的实测表现 |
|---|---|---|
| 普通交叉熵 | 类别均衡的基线对照 | macro AUC 最低,长尾类全崩 |
| 加权交叉熵 | 轻度不平衡 | 比基线好一些,参数敏感 |
| Focal Loss | 长尾明显、难样本多 | 涨点明显,$\gamma$ 需调 |
| 软标签 + KL | 有序分类 | 配合 Focal 使用效果最好 |
注意:改了损失函数之后,学习率通常需要重新调。Focal Loss 的梯度尺度比普通交叉熵大,我一般会把初始学习率降到原来的一半。
4. 训练流程与关键实现细节
4.1 优化器、学习率与训练轮次
优化器用 AdamW,权重衰减设 1e-4 到 5e-4 之间。为什么是 AdamW 而不是 Adam?因为 AdamW 把权重衰减从梯度更新里解耦出来了,正则效果更干净,小样本训练时这点差异挺明显。我对比过,用 Adam 的版本验证集波动更大,AdamW 更稳。
学习率初始值 1e-4,配合余弦退火加线性 warmup。warmup 的作用是训练初期不让主干被大梯度冲坏,特别是你用了预训练权重的时候。我的配置是前 3 个 epoch 从 1e-6 线性升到 1e-4,然后余弦退火到 1e-6。这个配置在几个不同主干上都跑得比较稳,可以直接抄。
训练轮次上,Batch Size 设 16 或 32(看你显存),总 epoch 设 50 到 80。小样本任务最怕的不是欠拟合而是过拟合,所以我建议盯着验证集指标做早停,连续 10 个 epoch 没提升就停,同时保存验证集上 macro AUC 最高的那个 checkpoint,而不是最后一个。
混合精度训练可以开,torch.cuda.amp在 EfficientNet 上大概能省 30% 到 40% 的显存,速度也快一点。要注意的是,开了 AMP 之后 BatchNorm 层偶尔会出现数值问题,如果 loss 突然变成 NaN,先把 AMP 关掉排查。
4.2 评估指标与验证策略
GAMMA 官方的评分指标我记得主要是 macro AUC,同时会看准确率和 Kappa。macro AUC 是先对每一类算 AUC 再取平均,对类别不平衡不敏感,这也是为什么长尾类崩了但普通准确率看着还行的时候,macro AUC 会直接暴露问题。
验证的时候有个容易出错的地方:五分类的 AUC 计算需要把标签做 one-hot,然后用每类的预测概率算。如果你直接用sklearn.metrics.roc_auc_score传多分类标签,得指定multi_class="ovr",不然会报错或者算出错误结果。我第一版代码就踩了这个坑,算出来的 AUC 明显偏高,后来发现是参数传错了。
验证阶段还有一个关键决策:模型选择用哪个指标。我建议用 macro AUC 和 Kappa 的组合,比如两个指标都看,取 macro AUC 最高且 Kappa 不低于峰值的 checkpoint。因为 AUC 高但 Kappa 低说明模型排序能力好但阈值划分差,实际提交可能会吃亏。
from sklearn.metrics import roc_auc_score, cohen_kappa_score import numpy as np def evaluate(model, loader, device, num_classes=5): model.eval() probs, gts = [], [] with torch.no_grad(): for batch in loader: cfp = batch["cfp"].to(device) oct_img = batch["oct"].to(device) logits = model(cfp, oct_img) p = torch.softmax(logits, dim=1) probs.append(p.cpu().numpy()) gts.append(batch["label"].numpy()) probs = np.concatenate(probs, axis=0) gts = np.concatenate(gts, axis=0) # 某些类别在验证集里可能一个样本都没有,需要跳过 present = [c for c in range(num_classes) if (gts == c).sum() > 0] auc = roc_auc_score(gts, probs[:, present], multi_class="ovr", average="macro", labels=present) preds = probs.argmax(axis=1) kappa = cohen_kappa_score(gts, preds, weights="quadratic") return auc, kappa, probs, gtsKappa 用二次加权(quadratic)是因为标签有序,预测偏一级和偏两级的惩罚应该不同。这个细节官方 Baseline 里不一定有,但加上之后更能反映真实性能。
4.3 完整训练脚本骨架
把上面的模块串起来,训练主循环大概是这样。
def train_one_epoch(model, loader, optimizer, scaler, criterion, device): model.train() total_loss = 0.0 for batch in loader: cfp = batch["cfp"].to(device, non_blocking=True) oct_img = batch["oct"].to(device, non_blocking=True) label = batch["label"].to(device, non_blocking=True) optimizer.zero_grad(set_to_none=True) with torch.cuda.amp.autocast(): logits = model(cfp, oct_img) loss = criterion(logits, label) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0) scaler.step(optimizer) scaler.update() total_loss += loss.item() * label.size(0) return total_loss / len(loader.dataset) def main(cfg): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") train_ds = GammaMultiModalDataset(train_df, cfg["train_root"], cfg, True) valid_ds = GammaMultiModalDataset(valid_df, cfg["valid_root"], cfg, False) train_loader = DataLoader(train_ds, batch_size=cfg["bs"], shuffle=True, num_workers=cfg["workers"], pin_memory=True) valid_loader = DataLoader(valid_ds, batch_size=cfg["bs"], shuffle=False, num_workers=cfg["workers"], pin_memory=True) model = GammaFusionNet(cfg).to(device) criterion = FocalLoss(gamma=2.0, weight=cfg["class_weight"]).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=cfg["lr"], weight_decay=cfg["wd"]) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_max=cfg["epochs"]) scaler = torch.cuda.amp.GradScaler() best_auc, patience = 0.0, 0 for epoch in range(cfg["epochs"]): tr_loss = train_one_epoch(model, train_loader, optimizer, scaler, criterion, device) auc, kappa, _, _ = evaluate(model, valid_loader, device) scheduler.step() print(f"epoch {epoch:03d} | loss {tr_loss:.4f} | " f"auc {auc:.4f} | kappa {kappa:.4f}") if auc > best_auc: best_auc, patience = auc, 0 torch.save(model.state_dict(), "best.pth") else: patience += 1 if patience >= cfg["early_stop"]: print("early stopping") break梯度裁剪那一步是我加的,官方 Baseline 里经常没有。小样本加上多模态融合,梯度偶尔会爆一下,裁剪到 5.0 能明显减少 loss 突然飞掉的情况。这个值不用太精细,5 和 10 之间差别不大。
5. 常见问题与排查速查
5.1 指标异常与数据问题
跑医学影像代码,指标出问题八成先怀疑数据,而不是模型。我整理了几个自己遇到过的典型症状和对应的排查路径。
| 症状 | 可能原因 | 排查方法 |
|---|---|---|
| loss 一直不降 | 标签对错位 / 学习率过大 | 打印一个 batch 的标签和图像,人工核对 |
| AUC 恒定 0.5 左右 | 输入全黑或全白 / 归一化错误 | 反归一化后把图存出来看一眼 |
| AUC 异常偏高 | 验证集泄漏 / 划分不当 | 检查是否用了官方划分,是否重复样本 |
| 训练集很好验证集崩 | 过拟合 / 域偏移 | 加增强、降模型容量、看混淆矩阵 |
| loss 变 NaN | AMP 数值问题 / 梯度爆炸 | 关 AMP,加梯度裁剪,降学习率 |
"输入全黑"这个问题我遇到过两次,都是归一化写错导致的。一次是 CFP 的像素值没除以 255,直接送进了Normalize,结果所有值都在 0 到 255 之间,减去均值之后分布完全错乱。另一次是 OCT 用了IMREAD_GRAYSCALE读成了单通道,然后np.stack复制三通道的时候维度顺序搞错了,存出来看是横条纹。凡是觉得模型没反应的,先把某个 batch 的第一张图反归一化存成 png,肉眼看看有没有问题,能省掉大量时间。
5.2 显存与训练速度
显存不够是最常见的工程问题。几个有效的降显存手段,按性价比排序:开混合精度(省 30% 到 40%)、减小 Batch Size(不够就配合梯度累积)、降低输入分辨率(从 512 降到 384 甚至 256)、换更小的主干。
这里有个折中要说明:降低分辨率会损失细节,而青光眼分级里杯盘比这类细粒度特征对分辨率敏感。我实测从 512 降到 256,macro AUC 掉大约 1.5 个点,如果显存实在紧张又不想掉分,可以只降 OCT 分支的分辨率,因为 OCT 的判别信息更集中在整体层结构上,对绝对分辨率不那么敏感。这个取舍我试下来比两路一起降要划算。
梯度累积的实现要注意 loss 要除以累积步数,否则等效于放大了学习率:
accum_steps = 4 for i, batch in enumerate(loader): loss = criterion(model(cfp, oct_img), label) / accum_steps scaler.scale(loss).backward() if (i + 1) % accum_steps == 0: scaler.step(optimizer) scaler.update() optimizer.zero_grad(set_to_none=True)训练速度方面,num_workers和pin_memory是基本盘,另外把数据预处理里能提前做的都提前做。比如尺寸统一、CLAHE、gamma 校正这些确定性操作,可以在数据准备阶段离线跑一遍存成缓存文件,训练时直接读,能省不少 CPU 时间。代价是占磁盘空间,几百例数据其实无所谓。
5.3 过拟合与跨中心泛化
小样本医学影像的过拟合几乎没有侥幸空间。我用的组合是:强数据增强、dropout、权重衰减、早停,四件套齐全。增强方面除了水平翻转,还可以加小角度旋转(±10 度以内)、随机亮度和对比度扰动、以及随机擦除。旋转角度不要太大,视盘的空间位置本身有解剖意义,转太狠反而破坏语义。
比过拟合更麻烦的是跨中心泛化。GAMMA 是多中心数据,不同中心的色温、分辨率、设备型号都不同,模型很容易学到"中心特征"而不是"病理特征"。我做过一个实验,把训练集按中心分组看验证集表现,发现某些中心的样本错误率明显高于其他中心,这就是域偏移的直接证据。
缓解手段上,颜色标准化比较有效:把 RGB 转到 LAB 空间,只对亮度通道做直方图匹配,把不同中心的图像对齐到一个参考分布上。这个方法实现不复杂,但对色温差异的鲁棒性提升明显。另外可以试试对抗式域适应,加一个中心分类的对抗头让特征更中心无关,但这个实现成本高,Baseline 阶段不推荐。
6. 提分实操:从 Baseline 到有竞争力的提交
6.1 数据增强与类别不平衡的组合拳
前面把增强和不平衡分开讲了,这里说说怎么组合。我的顺序是:先解决数据读取和划分正确性,这是地基;再加基础增强(翻转、缩放、亮度扰动)把过拟合压住;然后处理类别不平衡(加权 + Focal);最后做颜色标准化处理域偏移。这个顺序是有讲究的,前面一步没做稳,后面调参就是浪费时间。
类别不平衡的具体做法上,除了损失函数加权,还可以用重采样。但我不太推荐对少数类直接过采样,小样本下过采样容易让模型记住那几个样本,验证集看着好实际泛化差。更稳妥的是分层采样——让每个 batch 里的类别分布尽量均匀,同时保持整个 epoch 的样本数不变。实现上可以用WeightedRandomSampler,权重按类别频率的倒数设置。
注意:用
WeightedRandomSampler的时候,DataLoader的shuffle参数要设成 False,两者同时开会有冲突,PyTorch 会报错或者行为不符合预期。
6.2 模型集成与 TTA
单模型到瓶颈之后,集成是最稳的涨点手段。几个方向:不同随机种子的同结构模型集成、不同主干的集成、以及测试时增强。
多随机种子集成最容易实现,把训练脚本固定数据划分、只改随机种子跑三到五次,然后对测试集预测概率取平均。我实测五折或者五个种子集成,macro AUC 能稳定涨 2 到 3 个点,代价只是训练时间成倍增加。小样本任务里这个投入产出比很高。
测试时增强是对每张测试图做多种变换(原图、水平翻转、轻微旋转),分别预测后把概率平均。这个几乎不增加训练成本,涨点通常 0.5 到 1 个点。注意翻转 TTA 在多模态任务里要两路同步翻,不然空间对应关系就散了。
还有个容易被忽略的点:集成时如果各模型的预测概率分布尺度差异大,直接平均会被某个模型带偏。稳妥做法是先把每个模型的概率做一次温度缩放校准,再平均。温度参数可以在验证集上拟合,实现也就几行代码。
6.3 我踩过的坑与经验
最后说几个只在实际跑的时候才会暴露的问题。
第一个是图像和标签的对齐。我有一次图省事,用glob读文件列表然后按文件名排序去和标签表做 zip,结果标签表里的 ID 是字符串格式,图像文件名去掉了前导零,排序之后顺序全乱了。跑出来的结果全靠模型运气,训练 loss 看着还行但验证 AUC 死活上不去。后来加了一个断言,逐个比对 ID 才找出来。
第二个是验证集评估用错模式。有一次忘了写model.eval(),dropout 和 BatchNorm 还在训练模式,验证指标每次都不一样,来回震荡。调试了半天以为是学习率的问题,最后发现是这一行漏了。这种低级错误在赶进度的时候特别容易犯。
第三个是 O2O 数据命名不规范。有的中心 OCT 文件名带病例号的多个后缀,有的不带,硬编码模板路径会漏掉一部分样本。我的做法是在数据准备阶段先扫描一遍所有文件,把实际存在的文件名和标签表做个交集,生成一份干净的清单文件,训练时只读这份清单。这样即使原始数据命名再乱,训练阶段也只面对规范化的输入。
第四个是权限和路径问题。服务器上多人协作时,缓存的中间文件和 checkpoint 如果存在共享目录里,容易出现覆盖。我给每个实验配了一个独立的输出目录,目录名带上时间戳和关键超参,这样回头找实验记录的时候不用猜。
第五个是提交格式。GAMMA 的提交对文件格式、列名、样本顺序都有要求,我自己就在这上面浪费过一次提交机会——把预测概率提交上去了,而官方要的是类别标签。提交前一定拿官方给的样例文件对一遍列名和取值范围。
这套 Baseline 真正的价值不在它本身能跑多少分,而在于它把多模态医学影像分类的工程链路摆在你面前了。你把数据管线、融合模块、损失函数、验证策略这四块各改一处,就能体会到每一处改动对最终指标的影响,这比盲目调参有用得多。我个人在跑完三个完整实验周期后最大的感受是:先把数据和评估搞对,再谈模型,这句话在医学影像任务里怎么强调都不过分。