GAMMA青光眼分级Baseline精读:多模态融合与复现避坑
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任务一官方常用的评价指标是二次加权KappaQuadratic 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在加载预训练权重时往往会传出pretrainedTrue或者新版的weightsIMAGENET1K_V1然后把最后的全连接层替换成自己的分类头。替换的时候要注意输入维度——如果前面做了多模态拼接全连接层的输入维度是单模态特征维度 × 模态数这个数字算错了模型能定义成功但一前向就报维度不匹配。我在实操中养成的习惯是定义完模型立刻用一个假输入做一次前向把每个中间张量的shape打印出来。这一步能提前暴露90%的维度问题比等到训练报错再回头查效率高得多。提示torchvision不同版本的预训练权重API变化较大老代码里的pretrainedTrue在新版本里会报警告甚至失效复现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, modetrain): self.df df.reset_index(dropTrue) 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](imageimg)[image] images[modal] img label int(row[label]) return images, label这段代码有几个地方值得注意。第一reset_index(dropTrue)很关键因为如果你在划分数据集时用了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_classes5, backboneresnet34): 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(inplaceTrue), 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], dim1) 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-4SGD配1e-2到1e-3具体还要看batch size。optimizer torch.optim.Adam(model.parameters(), lr1e-4, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max50) 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 python3.8 -y conda activate gamma conda install pytorch1.10.0 torchvision0.11.0 cudatoolkit11.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×242M每个参数4字节约168MB。激活占用跟batch size成正比batch16时大约需要2-4GB。反向传播的梯度占用与参数同量级约168MB。加总下来batch16大概需要4-6GB显存batch32就需要8-12GB。我用过RTX 306012GBbatch开16比较稳开32偶尔会OOM。显存不够时优先降batch size其次考虑混合精度训练AMP后者能把显存占用降将近一半而且速度更快。学习率和batch size的关系遵循线性缩放原则——batch翻倍学习率大致也翻倍。但这不是铁律我一般把batch调小之后学习率也相应调小一点避免小batch下的梯度噪声太大导致训练不稳。这个规律在小数据集上尤其要谨慎套用。batch size建议初始学习率(Adam)显存需求(参考)85e-5约3GB161e-4约5GB322e-4约9GB643e-4约16GB4.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任务一上你大概率能少走我当年走过的那些弯路。