多模态眼底影像青光眼分级:GAMMA Baseline代码拆解与提分
发布时间:2026/9/30 1:39:10来源:尧图网络
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_trainTrue): self.df df.reset_index(dropTrue) 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(dropTrue) 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(clipLimit2.0, tileGridSize(8, 8)) return clahe.apply(img) def _gamma(self, img, gamma0.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, gamma0.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 到 8pin_memoryTrue训练时shuffleTrue验证时shuffleFalse。如果你用的是官方固定划分验证集的顺序也是固定的方便复现。3. 模型结构双分支编码器与融合方式3.1 主干网络选型Baseline 里两路通常共享同一个主干结构只是分别初始化。选型上要考虑三点数据量、显存、预训练权重可得性。GAMMA 的任务一训练样本只有几百例这个量级禁不起大模型折腾。我在实操里试过 ResNet-18、ResNet-50、EfficientNet-B0 三种。ResNet-50 参数量大小样本上过拟合明显验证集波动很大ResNet-18 稳定但天花板不高EfficientNet-B0 参数量约 5M配上预训练权重效果最好训练也快。最后我用的方案是两路各接一个 EfficientNet-B0CFP 分支输入三通道OCT 分支输入三通道灰度复制两路权重独立不共享。为什么不共享权重因为两路模态的统计分布差异太大共享编码器会强迫它们映射到同一特征空间反而拖累收敛。这个问题我做过对照实验共享权重的版本验证 macro AUC 比独立权重低 3 个点左右训练时长还多了将近一倍因为共享权重下梯度要同时兼顾两路收敛更慢。加载预训练权重的时候记得处理分类头不匹配的问题一般的写法是这样import torch.nn as nn from torchvision import models def build_backbone(nameefficientnet_b0, pretrainedTrue): if name efficientnet_b0: weights models.EfficientNet_B0_Weights.IMAGENET1K_V1 if pretrained else None net models.efficientnet_b0(weightsweights) 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(weightsweights) feat_dim net.fc.in_features net.fc nn.Identity() else: raise ValueError(funsupported backbone: {name}) return net, feat_dim这里把分类头换成Identity让主干直接吐特征向量后面统一接融合模块。这么做的好处是主干和融合层解耦你换主干的时候融合层代码一行都不用动。3.2 融合层从简单拼接到注意力这是整个 Baseline 里最值得动刀的地方。官方版本基本就是拼接可以写成这样class ConcatFusion(nn.Module): def __init__(self, feat_dim, num_classes5, dropout0.3): super().__init__() self.head nn.Sequential( nn.Linear(feat_dim * 2, 256), nn.BatchNorm1d(256), nn.ReLU(inplaceTrue), nn.Dropout(dropout), nn.Linear(256, num_classes), ) def forward(self, f_cfp, f_oct): return self.head(torch.cat([f_cfp, f_oct], dim1))拼接的假设是两路特征同等重要、彼此独立。但实际情况是OCT 在中晚期更关键CFP 在早期和整体形态上更关键两者权重应该随样本变化。所以我改成了门控式的注意力融合先算一个样本级的权重向量再用它加权求和。class GatedFusion(nn.Module): def __init__(self, feat_dim, num_classes5, dropout0.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(inplaceTrue), nn.Dropout(dropout), nn.Linear(256, num_classes), ) def forward(self, f_cfp, f_oct): g self.gate(torch.cat([f_cfp, f_oct], dim1)) # (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], dim1))这个门控的直觉是gate 输出接近 1 时模型更信任 CFP接近 0 时更信任 OCT。你可以把训练好的 gate 输出统计出来看分布我在验证集上观察过早期样本的 gate 均值确实偏大更依赖 CFP晚期样本偏小和临床认知是吻合的。这种可解释性在医学影像任务里挺有价值写报告的时候能拿来说事。如果你想再进一步可以做跨模态注意力把 CFP 特征当 QueryOCT 特征当 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_classovr不然会报错或者算出错误结果。我第一版代码就踩了这个坑算出来的 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_classes5): 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, dim1) probs.append(p.cpu().numpy()) gts.append(batch[label].numpy()) probs np.concatenate(probs, axis0) gts np.concatenate(gts, axis0) # 某些类别在验证集里可能一个样本都没有需要跳过 present [c for c in range(num_classes) if (gts c).sum() 0] auc roc_auc_score(gts, probs[:, present], multi_classovr, averagemacro, labelspresent) preds probs.argmax(axis1) kappa cohen_kappa_score(gts, preds, weightsquadratic) 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_blockingTrue) oct_img batch[oct].to(device, non_blockingTrue) label batch[label].to(device, non_blockingTrue) optimizer.zero_grad(set_to_noneTrue) 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_norm5.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_sizecfg[bs], shuffleTrue, num_workerscfg[workers], pin_memoryTrue) valid_loader DataLoader(valid_ds, batch_sizecfg[bs], shuffleFalse, num_workerscfg[workers], pin_memoryTrue) model GammaFusionNet(cfg).to(device) criterion FocalLoss(gamma2.0, weightcfg[class_weight]).to(device) optimizer torch.optim.AdamW(model.parameters(), lrcfg[lr], weight_decaycfg[wd]) scheduler torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_maxcfg[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(fepoch {epoch:03d} | loss {tr_loss:.4f} | fauc {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 变 NaNAMP 数值问题 / 梯度爆炸关 AMP加梯度裁剪降学习率输入全黑这个问题我遇到过两次都是归一化写错导致的。一次是 CFP 的像素值没除以 255直接送进了Normalize结果所有值都在 0 到 255 之间减去均值之后分布完全错乱。另一次是 OCT 用了IMREAD_GRAYSCALE读成了单通道然后np.stack复制三通道的时候维度顺序搞错了存出来看是横条纹。凡是觉得模型没反应的先把某个 batch 的第一张图反归一化存成 png肉眼看看有没有问题能省掉大量时间。5.2 显存与训练速度显存不够是最常见的工程问题。几个有效的降显存手段按性价比排序开混合精度省 30% 到 40%、减小 Batch Size不够就配合梯度累积、降低输入分辨率从 512 降到 384 甚至 256、换更小的主干。这里有个折中要说明降低分辨率会损失细节而青光眼分级里杯盘比这类细粒度特征对分辨率敏感。我实测从 512 降到 256macro 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_noneTrue)训练速度方面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 真正的价值不在它本身能跑多少分而在于它把多模态医学影像分类的工程链路摆在你面前了。你把数据管线、融合模块、损失函数、验证策略这四块各改一处就能体会到每一处改动对最终指标的影响这比盲目调参有用得多。我个人在跑完三个完整实验周期后最大的感受是先把数据和评估搞对再谈模型这句话在医学影像任务里怎么强调都不过分。
网站建设高端定制企业官网