尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

GAMMA青光眼分级多模态Baseline复现:CFP+OCT与QWK调优

GAMMA青光眼分级多模态Baseline复现:CFP+OCT与QWK调优 第一次把 GAMMA 任务一的官方 baseline 拉下来跑通我以为最耗时间的会是模型结构那几百行代码结果整整两天都耗在了 OCT 体数据的读取、显存分配和标签对齐上。这其实很典型MICCAI 2021 GAMMA 挑战赛的任务一叫“基于多模态眼底影像的青光眼分级”听起来是个四分类问题但它同时塞给你两种物理形态完全不同的输入——2D 的彩色眼底照CFP和 3D 的光学相干断层扫描体数据OCT。前者是二维图像后者是一个上百张切片叠起来的立方体。官方 baseline 给出的思路并不复杂复杂的是把它落到能复现、能出分、能改的状态。这篇就按我自己的复现顺序从分级口径、代码骨架、OCT 读取、调参记录一路拆到踩坑清单适合刚拿到这份 baseline 想跑通的同学也适合已经跑通但分数卡在瓶颈想往上抬的人。1. GAMMA任务一到底在考什么先分清分级口径再谈模型1.1 四档分级背后的临床判断逻辑很多人一上来就去看model.py这是最省事也最容易走偏的做法。分级任务的标签不是普通的“猫狗鸟”三类它对应的是青光眼视神经损害的程度分档通常从“无青光眼/非青光眼性改变”到“早期、中期、晚期”四档递进。也就是说这个四分类是有顺序的0 和 3 之间的距离远大于 0 和 1 之间的距离。理解这一点非常关键它直接决定了你应该用什么损失函数、用什么评价指标、以及为什么把 3 类错判成 0 类的惩罚要远大于错判成 2 类。青光眼的分级依据主要来自两块证据视盘的形态学改变视杯扩大、杯盘比增大、盘沿变窄、神经纤维层缺损和视网膜神经纤维层/神经节细胞层的厚度变化。CFP 擅长呈现前者——你一眼就能看到视盘的全貌和血管走行OCT 擅长呈现后者——它能给出黄斑区节细胞复合体和视盘周围神经纤维层的厚度分布。baseline 之所以要做多模态本质上就是因为这两块证据是互补的单看任何一个都会漏掉一部分病例。还有一个特别容易被忽略的点分级标签本身是医生综合判断的结果不同医生的判读之间存在一定的边界模糊。这意味着标签并非绝对干净靠近档位边界的样本天生就“难学”。你在训练后期看到验证集损失降不下去但准确率还在涨很可能不是过拟合而是模型在学那些边界样本的模糊区域。1.2 主指标为什么是二次加权Kappa而不是准确率GAMMA 任务一的官方主指标是二次加权 KappaQuadratic Weighted KappaQWK。我见过不少复现帖子直接用准确率做模型选择然后困惑于“为什么我准确率 0.7 提交上去分数不高”。原因很简单准确率把四档当成了四个平行类别错一档和错三档在它眼里完全一样而 QWK 会给错判加上按“档位距离平方”计算的惩罚权重。用公式说权重矩阵里相邻档位的惩罚是 1隔一档是 4隔两档是 9。这意味着“把重度判成中度”只扣一点点“把重度判成无青光眼”会扣得很惨。这个设计放在临床语境里是非常合理的——漏诊重度病例的代价确实远大于相邻档位的误判。所以你的整套训练策略都该围着 QWK 转模型选择看验证集 QWK早停看验证集 QWK后期调参也要看 QWK 而不是 loss。我在本地算 QWK 用的是 scikit-learnfrom sklearn.metrics import cohen_kappa_score # y_true / y_pred 均为 0~3 的整数标签 qwk cohen_kappa_score(y_true, y_pred, weightsquadratic)一个小提醒QWK 在验证集样本量只有百级的时候抖动会非常大。一个样本从 1 类挪到 2 类QWK 可能就跳 0.03。所以不要拿单次验证结果去判断模型优劣这也是后面我会建议做交叉验证的原因。1.3 CFP 与 OCT 在分级里各自承担什么角色如果把两种模态在分级任务里的贡献量化我的实测感受是CFP 提供了主要的“形态学先验”OCT 提供了主要的“厚度定量信息”两者叠加能显著改善中期档位的区分度。为什么改善最大的是中期档因为“无青光眼”和“晚期”这两头太好分了——一头的视盘完全正常另一头的杯盘比和视网膜神经纤维层厚度都到了肉眼可辨的程度单靠 CFP 就能判个八九不离十。真正难的是早期和中期之间的过渡带视盘形态改变还不明显但厚度已经开始掉了。这时候 OCT 的定量信息就成了决定性的那一票。反过来OCT 单独用也有明显短板。OCT 体数据的分辨率和成像范围限制了它对视盘整体形态的刻画而且不同设备的扫描协议、图像质量、分割算法差异会带来系统性偏差。所以 baseline 选择把两条分支并列而不是串联我认为是符合任务特性的串联会让一个模态的质量问题直接污染另一个模态的特征。2. 官方Baseline的代码骨架一次forward里发生了什么2.1 两条分支的backbone选型与输入张量形状官方 baseline 的整体结构是朴素的“双分支 后期融合”。CFP 分支把彩色眼底图当作标准 2D 图像处理走一个 2D 卷积主干baseline 用的是 ResNet 家族的常规配置输入张量形状是[B, 3, H, W]H 和 W 一般缩放到 256 或 512。OCT 分支拿到的是[B, D, H, W]的体数据其中 D 是切片数B-scan 的数量通常是上百这个量级。这里有个关键的设计取舍baseline 并没有直接把整段体数据丢进 3D 卷积网络而是做了切片层面的采样把体数据压成有限张切片再送进网络。原因很实在——3D 卷积在 D 上百、H/W 上千的输入上显存开销是灾难级的。你可以自己算一下一个[1, 1, 128, 512, 992]的 float32 体数据本身就是约 260MB只要主干里出现几个通道数翻倍的 3D 卷积层显存立刻爆掉。所以在读代码的时候你必须把注意力放在“它到底取了哪些切片、取了多少张、怎么拼成 batch”这几行上那不是无关紧要的预处理细节那是整个模型能不能跑起来的前提。2.2 融合层为什么用concat接全连接而不是注意力baseline 的融合层极其简单两条分支各自经过全局池化拿到一个特征向量然后torch.cat拼在一起过一层或多层全连接最后输出 4 维 logits。很多人会觉得这太“土”了为什么不用跨模态注意力、不用门控融合我的理解是baseline 的第一目标是“可复现、可跑通、无明显陷阱”而不是“刷最高分”。特征拼接的优点是无额外参数、不会因为某个模态质量差而产生梯度污染、调试时能清晰看到哪条分支在起作用。如果你要做消融实验拼接式融合也最容易做单模态退化。但它的缺点同样明显拼接不区分模态重要性模型只能靠全连接层的权重去隐式学习“什么时候该信 OCT、什么时候该信 CFP”。这在训练数据只有百级规模时很难学充分。所以如果你分数上不去融合层是最值得动刀的地方后面第 5 节我会展开讲几种改法。2.3 损失、优化器与训练调度baseline 的损失函数用的是标准的交叉熵优化器是 Adam学习率量级在 1e-4配合按 epoch 衰减的调度器。这套配置本身没什么特殊值得说的是几个容易出问题的地方。第一是学习率。如果 backbone 用的是 ImageNet 预训练权重1e-4 对新增的全连接层来说偏小对预训练主干又偏大。我自己的做法是分层设学习率主干小一点比如 1e-5 到 2e-5融合层和分类头大一点1e-4这样收敛快且不容易把预训练特征冲掉。第二是权重衰减和 dropout。百级样本量下过拟合几乎是必然的dropout 开在融合层比开在主干里更有效权重衰减我用 1e-4 到 5e-4 之间太大反而会让模型学不动 OCT 那种弱信号。第三是 batch size。因为 OCT 分支吃显存batch 往往被压到 2 或 4这时候 BatchNorm 的统计量会非常不稳定。我建议在这种小 batch 下把主干里的 BN 换成 GroupNorm或者至少把 BN 的 momentum 调小一点不然训练曲线会像心电图一样抖。2.4 推理阶段与提交文件的生成推理阶段 baseline 写得比较直白加载权重、遍历测试集、argmax拿到预测档位、写进 CSV。我在这块踩过两个坑值得单独说。第一个坑是推理时忘了切eval()模式。BN 和 dropout 在 train 模式下会引入随机性导致同一个样本跑两次结果不同而这种问题在代码里完全不会报错你只会在提交后发现分数莫名其妙地低于验证集。养成习惯训练循环里在每个 epoch 结束后加一次验证验证开头第一行就是model.eval()验证结束再model.train()。第二个坑是预测输出的顺序问题。测试集的样本顺序必须和 CSV 里的行顺序严格对应baseline 一般靠文件名排序来保证但如果你中途改过DataLoader的shuffle或者加过任何会打乱索引的采样器就一定要用文件名做 key 回填不要靠位置索引硬对。3. OCT体数据读取baseline里最容易被低估的一段代码3.1 先把显存账算清楚再动手我强烈建议在改任何模型结构之前先拿一段脚本把单个样本的体数据读出来打印它的 shape 和 dtype算一下内存占用。常规的 OCT 体数据规模可能是 D 在 100 到 200 之间、单张 B-scan 宽度接近 1000 像素。以[128, 512, 992]的 float32 为例单样本就是 128×512×992×4 字节约 260MB。这是“一个样本、一个通道”的原始开销。接下来做加法如果你要把体数据转成多通道比如复制成 3 通道喂给 2D 主干乘 3如果 batch 是 4再乘 4如果主干第一层就把通道数翻到 64中间特征图的开销会再上一个数量级。这样算下来一张 12GB 的卡在什么都不做的情况下就已经很紧张了。所以你会看到 baseline 里必有一步“降采样 切片采样”。这不是偷懒这是被硬件逼出来的工程妥协。我的经验是空间尺寸先降到 256 甚至 224这是性价比最高的一步切片数量控制在 16 到 32 张之间再往上加带来的收益很快就不明显了。3.2 切片采样策略的几种选择与实测差别baseline 里用什么策略取切片直接决定了 OCT 分支看到的信息量。常见的有这么几种策略做法优点风险等间隔采样从 D 张里均匀取 N 张覆盖整个体数据实现简单会跳过关键病灶层中心区域采样只取中间连续若干张黄斑/视盘中心通常在中段扫描偏位时直接丢信息投影合成对整段做最大值或均值投影把 3D 压成 2D省显存层间结构信息被抹平学习式选片用注意力加权各切片上限最高训练成本高小数据易过拟合我自己实测下来等间隔采样配合中心区域做加权即中心附近多取几张、两端少取是性价比最高的。原因在于青光眼相关的神经纤维层缺损虽然可能出现在多个位置但诊断信息最密集的区域还是集中在黄斑和视盘附近。纯等间隔的问题是你可能刚好把最厚和最薄的那几张切片跳过去了。如果你打算用投影合成一定要注意投影前的对齐。OCT 体数据在层与层之间可能存在轻微的位置漂移直接做最大值投影会把漂移放大成伪影模型很可能学到这些伪影而不是真正的厚度差异。3.3 预处理流水线的几个必须统一的细节预处理这一块baseline 通常写得比较简略但恰恰是最容易导致“本地跑得挺好、换台机器就崩”的地方。要统一的至少有这几项灰度归一化方式必须固定。OCT 是单通道灰度数据有人用(x - 0.5) / 0.5有人用 min-max 归一化到 [0,1]有人直接用 ImageNet 的均值和标准差。这三种在数值上差别很大训练和推理必须用同一套而且如果你加载了预训练权重用 ImageNet 统计量是更稳妥的选择。裁剪与填充策略必须固定。CFP 图像常常不是正方形直接 resize 会拉伸视盘改变杯盘比——这对青光眼分级来说是致命的形变。我一般用先按短边缩放到目标尺寸、再中心裁剪的方式保住几何比例。OCT 的层方向也要确认。不同设备的体数据在切片的排列顺序上可能不同有的从下往上有的从上往下。如果你把两边数据混着训模型会很困惑。建议先可视化几张切片确认方向一致。3.4 DataLoader的瓶颈与缓存技巧OCT 体数据的读取是典型的 IO 密集型任务。每次__getitem__都要从磁盘读一个几百 MB 的文件、解码、降采样、采样切片这个开销远大于 GPU 前向的时间。baseline 如果直接这么写你会看到 GPU 利用率长期在 30% 以下。解决办法有三个层次。最省事的是把num_workers调到 4 到 8pin_memoryTrue开启persistent_workers。中等成本的做法是把预处理后的体数据降采样 选好切片后的版本缓存成npy或npz文件第二次读就快得多代价是磁盘占用。最彻底的做法是用内存映射的方式直接读适合内存足够大的机器。我的实际配置是第一次跑预处理脚本把每个患者的体数据压成[N, 256, 256]的 npy 存下来训练时直接读 npy。这一改动让我的单 epoch 时间从 12 分钟降到了 3 分钟出头是最划算的一笔优化。4. 从零把Baseline跑到出分我的完整实操流程4.1 环境与依赖环境这块不复杂但版本要对齐。我用的是 Python 3.8 加 PyTorch 1.9 系列CUDA 11.1。医学影像相关的依赖主要用到opencv-python图像 IO 和 resize、scikit-image形态学处理和直方图均衡、scikit-learnQWK 和划分数据集、pandas标签文件读写、tqdm进度显示。有一个小细节值得提OpenCV 的cv2.imread默认是 BGR 顺序而大部分预训练权重是按 RGB 训练的。baseline 里如果没做cv2.cvtColor(img, cv2.COLOR_BGR2RGB)颜色通道就是反的。这个错误不会报错但会让预训练权重的效果打折。我自己第一次跑的时候就是因为这个验证集准确率比预期低了六七个点排查了很久才发现。4.2 数据清单与标签对齐这是整个流程里最枯燥也最不能出错的一步。我建议单独写一个校验脚本做三件事一是检查标签文件里出现的每个患者 ID在 CFP 目录和 OCT 目录里是否都能找到对应文件二是检查有没有文件存在但标签缺失的样本三是统计每个档位的样本数量。第三件事特别重要。青光眼分级数据的类别分布通常是不均衡的晚期和中期样本往往偏少。如果你的训练集里第 3 类只有个位数样本模型直接学到“永远预测 0 类”就能拿到不错的准确率但 QWK 会很差。这种情况下你必须处理类不平衡最直接的做法是加权交叉熵把类别权重按样本数倒数设置。对齐完之后把清单固化成一个 csv训练脚本只读这个 csv不要每次都在代码里去 glob 文件。这能避免“某次改了目录结构导致顺序变化”这种隐性 bug。4.3 关键超参设置与我用的值下面是我复现 baseline 时最终稳定下来的配置供参考参数取值说明输入尺寸(CFP)512×512保住视盘细节再大显存吃不消输入尺寸(OCT)256×256 × 24 切片切片数再多收益递减batch size4受 OCT 分支限制初始学习率主干 2e-5 / 头部 1e-4分层设置优化器Adam与 baseline 一致权重衰减2e-4中等强度正则损失加权交叉熵权重按类别频率倒数训练轮数60配合早停实际约 40 轮收敛数据增强随机水平翻转、轻微旋转(±10°)、亮度抖动不做强形变早停准则验证集 QWK 连续 8 轮不提升不用 loss 做准则关于数据增强我想强调一点眼底图不适合做激进的几何变换。垂直翻转会破坏上下方视网膜的解剖语义大角度旋转会让视盘位置产生不合理的偏移弹性形变更是会直接改变杯盘比的视觉判断。轻微的水平翻转因为左右眼本身就有镜像关系加小角度旋转基本就是安全边界。4.4 实测结果与每一档提升的来源我记录了几个阶段的结果能比较清楚地看出每一步改动值多少钱阶段配置验证集 QWK第一版纯 CFP 单分支不加权0.55 左右第二版加上 OCT 分支0.63 左右第三版加权交叉熵处理不平衡0.67 左右第四版分层学习率 预训练后冻结微调0.70 左右第五版5 折交叉验证 权重平均0.73 左右从这张表能看出最值钱的一步是“加上 OCT 分支”直接贡献了将近 8 个点这印证了多模态融合在分级任务上的必要性。第二值钱的是“5 折交叉验证加权重平均”它本质上是免费的分数——模型结构一点没变只是把 5 个模型的预测概率平均QWK 就上去了 3 个点。这在百级样本的任务里几乎是必做的操作。5. Baseline的短板与往上抬分的几条路5.1 先做单模态退化实验搞清楚信息分布在动手改结构之前我建议先做一件事把 OCT 分支关掉只用 CFP 训一遍再把 CFP 分支关掉只用 OCT 训一遍。这两个数字能告诉你很多信息。如果单 CFP 的 QWK 已经接近双模态说明 OCT 分支基本没学到东西问题很可能出在切片采样或归一化上而不是融合方式。如果单 OCT 明显低于单 CFP那你要考虑是不是 OCT 的预处理丢失了太多信息。我在排查时就发现过这个问题OCT 分支的验证 loss 一直在降但 QWK 不涨最后定位到是切片方向搞反了模型在学一个没有解剖意义的镜像结构。这个实验还有一个隐含价值它能帮你判断该往哪个方向投入。如果两个模态各自都不行优化融合是浪费时间如果两个单独都还行但融合不涨那就是融合方式的问题了。5.2 有序回归与类不平衡的联合处理前面提到四档分级是有序的但交叉熵完全无视这一结构。一个自然的改进是把问题转成有序回归训练三个二分类器分别判断“是否大于等于 1 档”“是否大于等于 2 档”“是否大于等于 3 档”推理时把三个概率累加得到最终的档位期望。这种做法的好处是它天然尊重档位顺序而且每个二分类器的正负样本比会比较均衡缓解了类别不平衡。代价是多了一层超参三个分类器的阈值不一定是 0.5你需要单独调。另一个思路是在损失函数里加一个“距离一致性”正则项惩罚相邻档位预测概率的剧烈跳变。这在只有百级样本时效果不太稳定我试过几次收益时有时无建议作为备选。5.3 融合层可以动的几个地方回到第 2 节提到的问题简单 concat 无法区分模态重要性。几种可以尝试的改法按实现难度从低到高排列最简单的做法是给两条分支的特征各过一个门控sigmoid 输出的标量权重加权后再拼接。这等价于让模型学一个自适应的模态置信度参数只多了几十个。中等做法是做双向交叉注意力用 CFP 的特征去 query OCT 的特征再用 OCT 的特征去 query CFP 的特征两个方向的结果拼接。这样两条分支在融合前就能互相“看一眼”对于“CFP 说视盘可疑、OCT 说厚度正常”这种冲突样本会更有分辨力。复杂做法是把 OCT 的切片特征序列当作时序用一个小型 Transformer 编码层做层间建模再和 CFP 特征做融合。这条路上限最高但在百级数据上极易过拟合必须配合强正则和大规模预训练我建议放在有其他数据集做预训练的前提下去尝试。5.4 交叉验证与集成的落地细节百级样本的任务里单次划分的验证结果参考价值很低。5 折交叉验证是基本操作但实现上有几个细节会影响最终收益。第一是划分方式要按患者而不是按图像。同一个患者可能有左右眼两张 CFP 和两个体数据如果按图像随机划分同一只眼睛的左右眼可能一个在训练集一个在验证集造成信息泄漏验证分数虚高。第二是分层抽样。划分时用StratifiedKFold按标签分层保证每折的类别分布一致。否则某一折可能完全没有第 3 类样本那一折的 QWK 就没有意义。第三是集成方式。最稳妥的是对预测概率取平均再 argmax而不是对档位投票。概率平均保留了模型的不确定性信息在有序指标下通常比投票好一截。如果想更进一步可以对各折模型验证集 QWK 做加权平均但要小心过拟合验证集——权重最好用固定值比如等权而不是搜出来的。6. 复现时我踩过的坑一条条列给你6.1 标签与文件名的隐性错位我遇到过一次非常隐蔽的问题CFP 的文件名里带左眼右眼标识OCT 的文件名里没有两者靠患者 ID 匹配。但有一批患者的 ID 在两个目录里的命名规则不一致一个带前导零一个不带匹配代码用了字符串精确比较结果这批样本被静默丢掉了。因为丢的正好是最难的那一档验证 QWK 看起来还挺高只有提交上去才发现分数掉得离谱。防范办法就是在数据校验脚本里打印匹配成功的样本数、各类别分布和人手数一遍的预期做对比。数量对不上就一定有问题。6.2 训练和推理的预处理不一致这是医学影像复现里最经典的一类 bug。训练时用了随机裁剪、随机亮度抖动推理时忘了关掉或者反过来训练时做了归一化但推理路径里漏了。表现是验证集正常但提交结果全乱而且极难排查因为代码不报错。我的做法是把预处理封装成一个类训练和推理共用同一个实例只用参数控制是否启用随机增强。这样物理上杜绝了两套代码不同步的可能。6.3 验证集剧烈抖动与早停失效小数据集加小 batch 的组合会让验证指标剧烈抖动。这时候如果你用“连续 N 轮不提升”做早停很可能在第 15 轮就停了而模型其实在第 30 轮才真正收敛到好的状态。两个缓解办法一是用滑动平均的验证指标做早停判断比如取最近 3 轮的中位数二是干脆不用早停固定训练轮数然后从保存的若干检查点里挑验证集最好的那个。我后来基本都走第二条路虽然多存几个模型但省心。6.4 提交格式与概率列的处理最后一个小坑提交 CSV 的列名、标签取值格式是 0-3 的整数还是字符串、是否需要同时提交概率每一届赛事的要求都可能不一样。baseline 里一般会带上写提交文件的代码但如果你自己改了推理流程很容易把这一环弄错。我的建议是拿一份官方的样例提交文件sample submission作为模板用同样的列名和 dtype 去写写完后立刻用 pandas 读回来assert一下形状和取值范围。花两分钟能省掉一次提交失败。我个人在实际操作中的体会是GAMMA 任务一这类多模态小样本竞赛任务真正拉开差距的不是模型结构有多花哨而是你有没有把数据对齐、预处理一致性、交叉验证这三件基础工作做扎实。我见过太多人在融合模块上折腾一周最后发现分数提升还不如把 OCT 切片方向改对来得实在。如果只让我给一条建议先花两天把数据管道做到闭着眼睛都能信再谈模型。
返回列表