1. 先搞清楚GAMMA任务一到底在考什么
GAMMA这个赛题,全称是Glaucoma Analysis with Multi-Modal imAging,是MICCAI 2021挂出来的多模态眼底影像分析挑战赛。它下面分了好几个任务,任务一(Task 1)聚焦的就是青光眼分级,也就是给一张或多张眼底影像,判断它属于青光眼的哪个严重等级。我最早接触这个赛题是冲着"多模态"三个字去的,因为当时手上正好有一批彩色眼底照和一批评估算相关的影像数据,想找个规范的benchmark来验证自己的融合思路会不会真的比单模态强。结果读完官方baseline代码之后发现,这套代码虽然结构不复杂,但它把多模态数据组织、分级任务标签定义、评价指标这几件最容易出错的事都处理得比较干净,非常适合拿来当自己项目的脚手架。
青光眼分级这个任务,说人话就是:医生看眼底图的时候,会重点关注视杯和视盘的相对大小,也就是所谓的杯盘比(CDR)。杯盘比越大,视神经被压迫得越厉害,青光眼就越严重。GAMMA任务一就是把这个判断过程变成一个多分类问题,通常按严重程度分成几个等级。你别小看这个"分级"和普通的"二分类"差别很大——二分类只要判断有病没病,而分级要求模型必须理解等级之间的序关系,把G3错判成G4,和把G3错判成G0,性质完全不一样。这就直接决定了评价指标的选择,后面我会专门讲。
1.1 青光眼分级任务的真实定义与标签体系
官方baseline里,标签是以整数形式给出的,属于有序分级(ordinal classification)。这一点在读代码时特别容易忽略:如果你把它当成普通的交叉熵多分类来做,理论上能跑,但会丢掉等级之间的有序信息,最终指标往往卡在某个瓶颈上上不去。我在复现的时候一开始就是这么干的,训练集上的loss降得很漂亮,验证集的Kappa却一直不理想,后来把标签的有序性显式地加进损失函数,才把差距补回来。
标签的取值范围通常是有限的几个等级,baseline里会用一个映射表把原始标注(可能是文本标签,也可能是某种临床分级)转成0到N-1的整数。你必须先搞清楚你手上的数据集的等级定义和官方是否一致,否则训练出来的模型换个数据集就完全没法用。我见过有人直接拿公开眼底数据集训练,结果等级划分标准跟GAMMA不一致,迁移过来预测全是乱的。所以第一步永远是核对标签体系。
1.2 多模态影像数据的组织方式与目录结构
GAMMA的多模态精髓在于,同一个病例可能对应多种成像方式,比如彩色眼底照(CFP)和光学相干断层扫描(OCT)。官方baseline一般会假设每张影像有唯一的文件名或ID,通过这个ID把不同模态的数据关联起来。这就意味着你的数据组织必须满足"一个样本ID对应多个模态文件路径"的映射关系。
我实际处理时,习惯先写一段脚本把目录结构打印出来,统计每个模态的文件数量、有没有缺失、命名规律是什么样的。这一步花不了十分钟,但能帮你省下几个小时的debug时间。baseline里通常用一个pandas的DataFrame来承载这个映射关系,每一行是一个样本,列里放不同模态的文件路径和标签。这种设计的好处是——你加一个模态只是加一列,代码改动极小;坏处是如果某个样本缺了某个模态,就得在Dataset里做特殊处理,不然读文件的时候直接报错。
注意:多模态数据最容易踩的坑就是"某个模态缺失",baseline往往假设数据完整,你自己接手真实数据时一定要加上缺失判断逻辑,否则训练跑到一半崩掉,找半天才发现是某张图丢了。
1.3 官方评价指标以及为什么用它
分级任务一般不单纯看准确率,因为它对序关系不敏感。GAMMA任务一官方常用的评价指标是二次加权Kappa(Quadratic Weighted Kappa,简称QWK)或者某个变体。QWK的核心思想是:预测等级和真实等级差得越远,惩罚越大,而且是按平方增长的。这正是为有序分级量身定做的。
为什么baseline里要专门实现这个指标而不是直接用sklearn默认的?因为多分类的Kappa实现里,加权矩阵的构造方式、标签的对齐顺序都可能影响结果。我在复现时就遇到过标签顺序不一致导致Kappa算出来是负数的尴尬情况——那其实不是模型烂,是评价代码写错了。所以读baseline的时候,评价指标那一段代码一定要逐行看懂,尤其是混淆矩阵怎么构造、权重矩阵怎么定义。
| 指标 | 适用场景 | 对序关系是否敏感 | 备注 |
|---|---|---|---|
| Accuracy | 一般分类 | 否 | 分级任务里容易虚高 |
| Macro F1 | 类别不均衡分类 | 否 | 各等级同等权重 |
| QWK | 有序分级 | 是 | GAMMA任务一常用 |
| AUC | 二分类/多分类 | 部分 | 需要额外处理多类 |
2. 官方Baseline的整体架构拆解
看完baseline之后,我最大的感受是:它没有追求花哨的融合结构,而是走了一条"特征提取 + 简单拼接 + 分类头"的稳妥路线。这种设计在竞赛baseline里很常见,因为它的首要目标是"保证能跑通、有合理基线分数",而不是刷到榜首。理解这一点很重要——不要把baseline当成标准答案去崇拜,而要把它当成一个可以放心魔改的起点。
从工程角度看,baseline把整个流程拆成了数据加载、模型定义、训练器、指标计算四块,耦合度比较低。我最喜欢的一点是它的数据集类写得比较通用,模态数量是可以通过配置调整的。这意味着你完全可以先拿单模态跑通,再逐步加模态,观察融合带来的增益,而不是一上来就被多模态的复杂度劝退。
2.1 为什么采用双分支多模态融合而不是单模态
多模态融合大体分三种:早期融合(early fusion,在输入层就把多模态拼起来)、中期融合(mid fusion,各自提特征后再融合)、晚期融合(late fusion,各自出预测再投票)。baseline一般选的是中期融合——每个模态走一个独立的特征提取分支,然后在高维特征层做拼接或相加。
为什么这么选?因为不同模态的数据分布差异巨大,彩色眼底照是二维的RGB图像,OCT是另一种纹理和对比度分布,直接在输入层拼通道,等于强行让一个卷积核去同时理解两种完全不同的视觉模式,效果往往不好。中期融合让每个分支先各自学到自己模态的表示,再在语义层融合,鲁棒性明显更高。我在自己的项目里对比过这两种方式,中期融合在小样本下优势尤其明显,验证集分数大概能高出几个百分点。
但中期融合也有代价——参数量翻倍、显存占用翻倍。如果你的显卡不够大,就得考虑共享backbone或者用更轻量的分支,这是后话,第4节会具体算。
2.2 Backbone选型与预训练权重加载逻辑
baseline通常用torchvision里现成的ResNet系列,比如ResNet34或ResNet50,加载ImageNet预训练权重。为什么要预训练?因为眼底影像数据集规模通常有限,从零训一个卷积网络很容易过拟合。ImageNet预训练提供了通用的低级纹理和边缘特征,迁移到医学影像上虽然不完美,但比随机初始化强太多。
这里有个细节值得说:baseline在加载预训练权重时,往往会传出pretrained=True或者新版的weights=IMAGENET1K_V1,然后把最后的全连接层替换成自己的分类头。替换的时候要注意输入维度——如果前面做了多模态拼接,全连接层的输入维度是单模态特征维度 × 模态数,这个数字算错了,模型能定义成功但一前向就报维度不匹配。
我在实操中养成的习惯是,定义完模型立刻用一个假输入做一次前向,把每个中间张量的shape打印出来。这一步能提前暴露90%的维度问题,比等到训练报错再回头查效率高得多。
提示:torchvision不同版本的预训练权重API变化较大,老代码里的
pretrained=True在新版本里会报警告甚至失效,复现baseline时务必先确认你的torchvision版本与代码匹配。
2.3 数据增强与预处理管线的设计思路
眼底影像的增强有讲究。普通的随机裁剪、翻转可以用,但要注意左右眼是对称的——水平翻转会把左眼变成右眼,如果你的标签或后续特征里包含了眼别信息,翻转就会引入错误。baseline里一般只做基础的几何增强和归一化,这是稳妥做法。
预处理里最关键的是归一化参数。医学影像的像素分布跟自然图像差别很大,如果你直接用ImageNet的均值和方差做归一化,虽然能用,但不一定最优。baseline用ImageNet统计量,是为了配合预训练权重。如果你打算从零训练,不妨统计一下自己数据集的均值和方差,通常能带来一点提升。
还有一点:不同模态的预处理管线应该是独立的,因为它们的像素值范围、色彩空间可能完全不同。baseline在Dataset里给每个模态配一套transform,这个设计是对的,你在魔改时千万别图省事把所有模态塞进同一个transform。
3. 核心代码逐块精读与实现细节
这一节我打算把baseline里最关键的几段代码拎出来,配上我的解读和实操注释。代码本身不长,但每一行背后都有它的道理,看懂了这些,你魔改起来才不会心虚。
3.1 Dataset与DataLoader的模态对齐实现
Dataset类是整个数据管线的入口。baseline的写法一般是这样:初始化时传入一个DataFrame,里面每行是一个样本,列里放各模态的路径和标签;__getitem__里逐个读取模态文件,做transform,然后打包成一个字典或元组返回。
class GammaDataset(Dataset): def __init__(self, df, transform_dict, mode='train'): self.df = df.reset_index(drop=True) self.transform_dict = transform_dict # 每个模态一套transform self.mode = mode def __len__(self): return len(self.df) def __getitem__(self, idx): row = self.df.iloc[idx] images = {} for modal in self.modal_list: # 例如 ['cfp', 'oct'] img = cv2.imread(row[f'{modal}_path']) img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img = self.transform_dict[modal](image=img)['image'] images[modal] = img label = int(row['label']) return images, label这段代码有几个地方值得注意。第一,reset_index(drop=True)很关键,因为如果你在划分数据集时用了df.sample(),索引会乱掉,iloc和loc混用就会出错。我就被这个坑过一次,训练集和验证集数据串了,模型表现异常好,排查半天才发现是索引问题。第二,cv2读进来是BGR格式,一定要转RGB,否则颜色通道反了,预训练权重的效果大打折扣。第三,返回的是字典结构,这样在模型里可以按模态名取用,比硬编码位置更灵活。
DataLoader部分没什么玄机,设好batch_size、num_workers、shuffle就行。num_workers设大一点能加快数据加载,但设太大反而会因为进程切换拖慢速度,一般设成CPU核心数的一半比较稳。我在8核机器上一般设4,实测下来比较平衡。
3.2 多模态特征融合层的代码实现
模型的主体是双分支加融合。backbone负责提特征,融合层负责在特征维度上做拼接,最后接分类头。
class MultiModalNet(nn.Module): def __init__(self, num_classes=5, backbone='resnet34'): super().__init__() self.branch_cfp = self._build_backbone(backbone) self.branch_oct = self._build_backbone(backbone) feat_dim = self._get_feat_dim(backbone) # 例如 resnet34 是 512 self.classifier = nn.Sequential( nn.Linear(feat_dim * 2, 256), nn.ReLU(inplace=True), nn.Dropout(0.5), nn.Linear(256, num_classes) ) def forward(self, images): f_cfp = self.branch_cfp(images['cfp']) f_oct = self.branch_oct(images['oct']) fused = torch.cat([f_cfp, f_oct], dim=1) return self.classifier(fused)融合用torch.cat是最直接的。有些baseline会用加法或注意力加权,但拼接是最稳的,因为它保留了两个模态的全部信息,让后面的全连接层自己去学怎么权衡。Dropout放在融合后很重要,因为多模态拼接后特征维度变高,过拟合风险随之增加。我在自己项目里试过把Dropout去掉,训练集准确率飙到接近100%,验证集直接不动了,典型的过拟合。
还有个易错点:feat_dim的取值。ResNet34和ResNet50的最后一层特征维度都是512,但如果你换了backbone,比如换成EfficientNet,就得重新确认。写死数字是大忌,最好用一个小函数动态获取,避免换backbone时忘了改。
3.3 损失函数、优化器与学习率调度配置
baseline默认用交叉熵损失,优化器用Adam或SGD加动量。这里我强烈建议你根据自己的数据做调整。如果各等级样本数量严重不均衡(青光眼中重度样本通常很少),交叉熵会偏向多数类,你应该上加权交叉熵或者Focal Loss。我实测过,加类别权重后,少数类的召回率能明显改善,整体Kappa也更高。
学习率调度baseline一般用StepLR或者CosineAnnealingLR。前者简单,到设定epoch就衰减一次;后者平滑,但需要你给出总epoch数。我个人的偏好是先用CosineAnnealingLR跑一版看曲线,如果收敛不稳再退回StepLR。初始学习率设太大,loss会震荡;设太小,收敛慢得像蜗牛。经验值是Adam配3e-4到1e-4,SGD配1e-2到1e-3,具体还要看batch size。
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=50) criterion = nn.CrossEntropyLoss()weight_decay别忽略,它相当于L2正则,对抑制过拟合有帮助。但注意,如果你用了带权重衰减的AdamW,那又是另一回事了,AdamW把权重衰减和优化分离,一般效果更好,新项目里我会优先选它。
3.4 训练与验证循环的完整实现
训练循环是模板化的,但细节决定成败。baseline里训练和验证通常写成两个函数,训练时切train()模式,验证时切eval()并且用torch.no_grad()包住。
def train_one_epoch(model, loader, optimizer, criterion, device): model.train() total_loss = 0 for images, labels in loader: images = {k: v.to(device) for k, v in images.items()} labels = labels.to(device) optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() total_loss += loss.item() return total_loss / len(loader)验证阶段要收集所有预测和标签,最后统一算QWK,不能一个batch一个batch地算再平均,那样得到的不是全局Kappa。这是个高频错误,我见过不少人在这里栽跟头。正确做法是把每个batch的预测累积到一个列表里,epoch结束后一次性算指标。
另外,模型保存策略也要注意:不要只存最后一个epoch的权重,要按验证指标存最优的。baseline一般会记录best_score,每次验证超过就覆盖保存。我吃过亏,中途某次指标最好但没保存,最后模型反而变差了,只能重训。
4. 从零跑通Baseline的完整实操流程
理论讲完了,这一节上硬货。我会把从环境配置到训练跑通的每一步都拆开,包括参数怎么算、日志怎么看。你可以直接照着抄。
4.1 环境准备与依赖版本选择
baseline一般基于PyTorch,版本匹配是头等大事。torch、torchvision、CUDA三者的版本必须对应,否则轻则警告重则直接报错。我的建议是去PyTorch官网的版本对应表确认,然后优先用conda装,因为它会自动帮你解决CUDA依赖。
conda create -n gamma python=3.8 -y conda activate gamma conda install pytorch==1.10.0 torchvision==0.11.0 cudatoolkit=11.3 -c pytorch -y pip install opencv-python pandas scikit-learn tqdmpip install opencv-python pandas scikit-learn tqdm为什么选Python 3.8?因为它是那一批baseline代码兼容性最好的版本,太新的Python有些老库会编译不过。cudatoolkit选11.3是因为对应的显卡驱动兼容范围广。你如果显卡是30系以后的新架构,可以适当往上选,但要注意别选超过你驱动支持上限的版本。
我踩过一个坑:装完发现torch.cuda.is_available()返回False,排查半天是conda装的版本和系统级pip装的版本冲突了。判断方法很简单,运行python -c "import torch; print(torch.__version__)"看看到底用的是哪个。环境干净比版本新更重要。
4.2 显存、batch size与学习率的参数权衡计算
显存占用可以粗略估算。以ResNet34双分支、输入224×224为例:
- 参数占用:两个ResNet34约21M×2=42M,每个参数4字节,约168MB。
- 激活占用:跟batch size成正比,batch=16时大约需要2-4GB。
- 反向传播的梯度占用与参数同量级,约168MB。
加总下来,batch=16大概需要4-6GB显存,batch=32就需要8-12GB。我用过RTX 3060(12GB),batch开16比较稳,开32偶尔会OOM。显存不够时,优先降batch size,其次考虑混合精度训练(AMP),后者能把显存占用降将近一半,而且速度更快。
学习率和batch size的关系遵循线性缩放原则——batch翻倍,学习率大致也翻倍。但这不是铁律,我一般把batch调小之后,学习率也相应调小一点,避免小batch下的梯度噪声太大导致训练不稳。这个规律在小数据集上尤其要谨慎套用。
| batch size | 建议初始学习率(Adam) | 显存需求(参考) |
|---|---|---|
| 8 | 5e-5 | 约3GB |
| 16 | 1e-4 | 约5GB |
| 32 | 2e-4 | 约9GB |
| 64 | 3e-4 | 约16GB |
4.3 训练日志怎么看与模型保存策略
训练日志里我最关注三样东西:训练loss、验证loss、验证指标。如果训练loss持续下降但验证loss开始上升,那就是过拟合的典型信号,该上早停(early stopping)了。如果两者都下不去,可能是学习率太小或者模型容量不够。如果loss剧烈震荡,多半是学习率太大或者batch太小。
指标方面,QWK在验证集上的波动往往比loss大,因为它对少数类的预测很敏感。我建议观察连续几个epoch的移动平均,别因为单次波动就急着调参。模型保存至少要保存两个文件:一个是按验证指标最优的权重,一个是最后一个epoch的权重(用于复盘)。
我习惯在训练脚本里加一段自动生成训练曲线的代码,把loss和指标画出来存成图片。跑了几十次实验之后你会发现,有一张曲线图比一堆数字好理解太多,回头找规律的时候也方便。
5. 训练过程中那些坑与排查技巧
baseline代码能跑通是一回事,能跑出好结果是另一回事。下面这些坑我是真金白银踩出来的,整理成速查表,你照着排查能省不少时间。
5.1 数据侧常见问题速查
数据侧的问题最隐蔽,因为程序不会报错,但结果就是不对劲。最常见的几个:标签错位(数据集划分后索引没重置)、模态不匹配(同一ID对应的两个模态其实是不同病例)、图像损坏(个别图读出来是全黑或全白的)。我强烈建议在训练前跑一遍数据体检脚本:统计每个等级样本数、检查每张图的尺寸和像素均值、验证模态ID能一一对应。这几步加起来也就几十行代码,但能提前拦下大部分脏数据问题。
还有一个容易忽略的点:青光眼分级里等级分布往往很不均衡,轻度样本一大堆,重度样本寥寥无几。如果你不做任何处理直接按原始分布训练,模型会倾向于全预测多数类,指标看起来还行其实完全没用。解决方式有重采样、类别加权、Focal Loss等,我一般先试类别加权,改动最小见效最快。
5.2 模型侧与训练侧常见问题速查
模型侧最典型的就是维度对不上。多模态拼接后全连接层输入维度算错,会在第一次前向时报错。解决办法是定义完模型后立刻用假数据跑一次前向,把shape打印出来核对。另外,如果多模态里某个模态的数据分布和预训练权重的假设差太多(比如单通道的OCT直接喂给三通道的backbone),要么复制通道,要么改第一层卷积,别硬塞。
训练侧最常见的是loss变NaN。原因通常是学习率太大、数据里有NaN、或者用了不稳定的损失函数。排查顺序是:先把学习率调小十分之一试试,不行再检查数据是否含异常值,最后看损失函数实现。我遇到过因为标注里有-1(表示未知)导致交叉熵计算出NaN的情况,这种脏标签必须提前过滤。
注意:如果验证集指标一直不动,先别急着改模型,回头看看学习率、标签、数据增强这三样,八成问题出在这里。
5.3 独家避坑经验
分享几条文档里不会写的经验。第一,训练初期先用一个极小的子集(比如100张图)跑通全流程,确认代码能端到端跑起来、能保存模型、能算指标,再上全量数据。全量数据跑一次可能几小时,小数据集几分钟就能验证代码正确性。第二,固定随机种子,否则你没法复现自己的实验,调参就变成了玄学。种子固定后,同样的代码同样的数据,结果应该完全一致。第三,多模态模型别一上来两个分支都用大backbone,可以先一个分支大、一个分支小,或者共享backbone,观察效果再决定要不要加容量。资源有限时,聪明的架构选择比堆算力更有效。
还有一条:验证集和测试集的预处理必须完全一致。有人训练时用了数据增强,验证时忘了关,结果验证指标虚高,测试时原形毕露。DataLoader在验证模式下一定要只做确定性的归一化,把所有随机增强关掉。
6. Baseline之后还能怎么往上做
baseline只是起点。如果你跑通了它、复现了它的分数,接下来就可以折腾一些提升方向了。这一节聊聊我自己试过或者觉得可行的思路。
6.1 多模态融合的进阶玩法
拼接是最基础的融合。往上可以做注意力融合——让模型自己学两个模态特征的权重,哪个模态对当前样本更重要就多听谁的。这在某个模态质量不稳定时特别有用,比如有些眼底照拍得模糊,模型就应该更依赖OCT分支。还有一种做法是跨模态注意力,让一个模态的特征去"查询"另一个模态的特征,捕捉模态之间的相关性。
再进一步,可以考虑每个模态单独预训练,再联合微调。baseline一般是用同一个ImageNet权重初始化两个分支,如果某个模态有自己的大规模预训练数据(比如眼底影像领域的自监督预训练模型),用它来初始化对应分支,往往能带来明显的提升。这是当前多模态领域比较主流的一个思路:先各自强,再融合强。
6.2 分级任务特有的技巧
分级不同于普通分类,充分利用标签的序关系能白捡一些分数。做法之一是用有序回归思路,把多分类问题拆成一串二分类(G0 vs 其余、G0-G1 vs 其余……),每个二分类器输出一个累积概率,最后推导出等级。这样模型的输出天然满足单调性,不会出现"预测为G3但累积概率却低于G2"这种矛盾。
另一个技巧是在损失函数里显式加入等级距离惩罚。预测和真实标签差得越远,惩罚越大,这跟QWK的加权思路是一致的。我试过在交叉熵基础上加一个距离项,验证集QWK提升了一点点,虽然不多,但比较稳定。还有,分级任务里少数类样本极其宝贵,可以考虑对少数类做更强的增强,甚至用生成模型合成一些,但要小心合成数据的真实性,别引入分布偏移。
换backbone、加预训练、调增强这些通用手段当然也能用,但我建议一次只改一个变量,做好记录,否则你永远不知道是哪个改动起了作用。这个习惯是我做竞赛时被队友逼出来的,后来发现它其实是从业者做实验的基本素养。
我个人在实际复现这类baseline时的体会是,它最大的价值不在分数,而在提供了一个结构清晰、边界明确的起点。你先老老实实把它跑通、看懂每一行,再带着问题去改,比一上来就魔改网络结构要高效得多。我通常会把baseline的配置、数据组织、评价代码三部分单独抽出来,作为后续所有实验的公共基础设施,这样每次尝试新想法时,改动都能控制在很小范围内。分享一个我常用的小习惯:每跑完一组实验,我会在笔记本里记下这次的配置、指标和一条"下次要注意什么",攒上几十条之后,再遇到问题基本能秒定位。这套方法用在GAMMA任务一上,你大概率能少走我当年走过的那些弯路。