多标签图像分类实战:HPA竞赛冠军方案全解析
发布时间:2026/10/2 15:21:14来源:尧图网络
简介这是一份Kaggle人类蛋白质图谱图像分类第一名解决方案的完整技术文档面向深度学习与计算机视觉竞赛爱好者重点解决多标签分类、类别高度不均衡、训练集与测试集分布不一致以及阈值调整与后处理调优等实际难题。压缩包仅含1个docx文档体积仅117KB便于快速下载阅读。在当前已有116人学习的基础上这份文档凭借紧凑的篇幅呈现了冠军方案的完整技术栈。内容系统复现了基于DenseNet121的CNN分类器结构包括AdaptiveConcatPool2d、Flatten、BatchNorm1d、Dropout、Linear与ReLU等模块的组合并详细说明了Adam优化器、分段衰减学习率、FocalLossLovasz损失函数的设计动机。数据方面通过哈希方法剔除v18外部数据中约6000个重复样本并采用traintest统计均值方差进行归一化同时使用旋转、翻转与随机裁剪增强训练数据。后处理阶段则利用保持标签比例和度量学习近邻修正两种提交策略有效提升了最终分数。文档还附有训练验证思路与竞赛心得兼具方法论和实战细节适合希望提升医学图像分类与Kaggle竞赛能力的读者。1. 人类蛋白质图谱竞赛为什么第一名方案值得你复现Kaggle人类蛋白质图谱图像分类赛题官方名称是Human Protein Atlas Image Classification任务是把显微镜下的人类细胞图像按蛋白质在细胞内的定位方式分类。每张图不是一张普通照片而是来自Human Protein Atlas项目的高通量染色图像需要同时预测多个标签比如蛋白质是否在细胞核、线粒体、内质网、高尔基体里最多可以有28个标签同时存在。第一次看到这个题的人容易把它当成普通的图像分类任务直接拿ImageNet预训练模型去微调结果往往在验证集上自嗨一上公共榜就现原形。2018年这届竞赛的第一名解决方案真正拉开差距的并不是model架构本身而是多标签数据清洗、损失函数设计、交叉验证和模型集成这一整条流水线。这篇文章不是复述某个不可考的“冠军源码”而是把这类显微镜图像分类竞赛公认的最稳妥打法拆成能落地的步骤。适合刚入门Kaggle竞赛但想挑战多标签图像的选手也适合做病理、荧光显微图像识别但一直觉得模型不收敛的从业者。你会看到为什么图像分类模型要用预训练CNN而不是直接上transformer也会看到多标签场景下BCE损失比交叉熵更合适以及TTA和集成在这里能带来多少可靠的收益。先把数据弄明白再谈模型。2. 读懂HPA图像数据三通道、28标签与不均衡样本的处理2.1 图像的三个通道到底在拍什么HPA赛题提供的是灰度样式的细胞图像但实际以RGB三通道存储每个通道对应一次染色标记。红通道通常是微管或某种细胞器蛋白绿通道是目标蛋白质本身的染色蓝通道是细胞核的DNA染色。三通道合在一起看能同时告诉你“细胞轮廓在哪”“蛋白质分布在哪”“核在哪”。做预处理时不要把这三通道当成自然图像的RGB否则模型学到的颜色语义完全不对。常见做法是把每个通道独立做对比度归一化然后合回三通道。我一般不用自然图像那套RGB均值方差归一化而是先统计每个通道的像素强度分布做直方图裁剪。比如把1%分位以下压到099%分位以上压到255再用均一化。这一步能显著提升后续模型的收敛速度因为显微图像的动态范围经常很大不裁剪的话训练时模型会被高亮噪声带偏。处理完图像后你会面对一个更现实的问题类别极度不均衡。28个标签里最常见的是“细胞质”“细胞核”“微管蛋白”每个出现几千次而像“紧密连接”“核仁纤维中心”这类标签整个训练集可能只有几十张。直接训练会学到“全部预测0”这种无聊的答案。解决不均衡有两种主流思路一是标签加权loss把低频标签的损失权重调高二是做分层采样或人工扩充。第一种更常用因为实现简单且不会引入伪造样本。2.2 多标签分类的标签编码与样本权重计算多标签分类要求每个样本同时属于多个类别标签不是one-hot而是multi-hot向量。HPA竞赛给的是每幅图像对应的一组标签ID你需要把ID映射到027的索引然后用长度28的0/1向量表示。这步做错的人很多常见的错误是把标签当成单标签去算softmax交叉熵。在PyTorch里加载标签的代码骨架如下import pandas as pd import numpy as np from PIL import Image label_names sorted(set(...)) # 从训练集里收集全部标签 label_to_idx {lab: i for i, lab in enumerate(label_names)} num_classes len(label_names) def read_sample(img_path, id_to_labels, label_to_idx): img Image.open(img_path).convert(RGB) labels id_to_labels[id] # 这个id对应的标签列表比如[Nucleus,Cytoplasm] vec np.zeros(num_classes, dtypenp.float32) for lab in labels: vec[label_to_idx[lab]] 1.0 return img, vec这段代码的核心是把每个样本的标签集合展开成长度固定的向量。注意vec必须用float32而不是int因为后续要交给BCEWithLogitsLoss计算它内部要处理logits与target的形态。用int会带来一些莫名其妙的反向传播报错。样本权重的计算直接决定低频标签能不能被模型记住。一个有效做法是对每个标签统计其在训练集出现的样本数然后取倒数的对数或开方形成该标签的权重。再把这个标签权重用在BCEWithLogitsLoss的pos_weight参数上。from torch.nn import BCEWithLogitsLoss # label_freq: 长度为 num_classes 的数组每类出现次数 pos_weight (1.0 / label_freq) ** 0.5 pos_weight pos_weight / pos_weight.mean() # 归一化避免整体梯度过大 criterion BCEWithLogitsLoss(pos_weighttorch.tensor(pos_weight, devicedevice))这里pos_weight的含义是对于正类样本的损失乘以一个放大系数。低频标签因为出现次数少它的1 / freq会偏大于是模型更倾向于把这些样本预测为存在。取0.5次方是为了防止权重差异过大导致训练震荡。如果你发现低频标签学了但过拟合可以在0.30.7之间调这个指数。2.3 数据清洗与折叠划分先做预处理再碰模型HPA数据里存在不少标错、模糊或背景过亮的图像。第一名方案通常会做一轮手动或半自动清洗最常见的是删除“完全没有染上色的空白图”以及标签明显矛盾的样本。简单规则可以看每个通道的像素方差如果三通道方差都特别低基本就是空白或失焦图可以直接从训练集中剔除。折叠划分不要直接随机切。多标签场景下最好按“样本包含的标签组合”做分层保证每一折里低频标签的分布和全集一致。实现上可以给每个样本生成一个多标签指纹比如把所有标签拼接成一个字符串然后用这个指纹做分层KFold。这样做比完全随机划分更稳因为随机划分可能导致某一折里特别稀有的标签一个都没出现模型在那折验证会低估性能。划分代码示意如下from iterstrat.ml_stratifiers import MultilabelStratifiedKFold mlkf MultilabelStratifiedKFold(n_splits5, shuffleTrue, random_state42) for fold, (tr_idx, va_idx) in enumerate(mlkf.split(X_ids, y_multi_hot)): # X_ids 是样本id数组y_multi_hot 是(n_samples, num_classes)的one-hot矩阵 print(ffold {fold}: train {len(tr_idx)}, valid {len(va_idx)})这个库在PyTorch生态里很常用。如果不想额外依赖也可以自己写一个基于标签频率的分层分组逻辑但MultilabelStratifiedKFold已经是社区里最成熟的方案。折叠数一般5折就够既保证验证置信度又不至于把训练时间拉太长。如果你用transformer架构做对比折叠数可能要减少因为训练成本高但基线CNN方案5折很从容。3. 用预训练CNN做图像分类模型选型、损失与训练配置3.1 为什么第一梯队方案不是从零训练的transformerHPA图像属于显微图像尺寸通常只有512x512或类似目标在图像上的尺度较小而且图像染色后的纹理模式与自然图像差异很大。竞赛结束那几年最稳的模型是预训练CNNResNet、Inception、DenseNet、SE-ResNeXt之类在ImageNet上的权重再做迁移学习。近年来transformer图像分类模型ViT、DeiT、Swin确实在自然图像上很猛但在小而专的显微图像上如果数据量只有几万张纯transformer从零训练或者只做最后几层微调往往不如卷积模型加充分的数据增强来得稳定。原因很简单transformer需要大量数据学到归纳偏置而HPA训练集规模有限CNN预训练权重已经含了很多纹理、边缘、形状的低级特征这些特征在显微镜图像里同样有效。第一名方案的主力基本都是强化版的ResNet或Inception类网络再配各种数据增强。如果你非要用transformer建议作为一个额外的集成个体比如用Swin Tiny在ImageNet上预训练后做全量微调放进去当“多样性”来源而不是指望它单独拿第一。3.2 用PyTorch加载预训练模型并改造多标签输出层预训练模型从torchvision.models加载最简单。改动的地方只有最后一层把输出类别数从1000改成28。因为BCEWithLogitsLoss接受logits所以最后一层不要接Sigmoid直接输出28维原始数值。import torchvision.models as models model models.resnet50(weightsmodels.ResNet50_Weights.IMAGENET1K_V1) fc_in model.fc.in_features model.fc torch.nn.Linear(fc_in, num_classes) # num_classes28如果你的图像尺寸大且显存吃紧可以把第一个7x7卷积的kernel改成3x3或把stride改成1用来适配小图输入。但这样做会丢掉一部分预训练权重对应的感受野你要在分辨率调低和模型结构改动之间做取舍。常见做法是保持默认resnet结构只把输入图片resize到256x256或384x384而不是改主干。实际上第一名方案往往用更大的backbone比如SE-ResNeXt50或Inception ResNet v2。但这些模型在torchvision里不一定都有预训练权重有时需要自己加载第三方权重。为了可复现性我建议先从ResNet50和Se-ResNet50开始跑通蒸馏流程后再换更大模型。3.3 损失函数选择BCEWithLogitsLoss和类别加权多标签图像分类的默认损失是二分类交叉熵的总和PyTorch里就是BCEWithLogitsLoss它把Sigmoid和BCE融合在一起数值上比手动分开算更稳定。为什么不直接用多标签softmax因为softmax强制所有标签互斥而蛋白质在多个细胞器同时现身是很常见的softmax会逼迫模型在竞争关系里选一个最终损失很大但精度上不去。我在训练时会把损失拆成两部分来看总损失是28个标签的BCE之和但每个标签对总损失的贡献不同。为了监控每个标签的学习我会在每个epoch结束后单独计算每个标签的AUC再取micro平均。如果某个低频标签的AUC一直低于0.7就需要加重它的pos_weight。这个操作可以直接反映在验证指标上比盯着全局loss直观得多。BCEWithLogitsLoss还有一个容易被忽略的参数reduction。我之前习惯设成sum结果是梯度大小被batch size放大学习率稍微调大就发散。后来改成mean损失被batch内所有像素类别平均训练稳定很多。对多标签场景更稳的是把pos_weight乘上后再做mean。如果你发现验证曲线震荡先检查这个参数。3.4 训练超参batch size、lr schedule与混合精度训练HPA图像分类我不建议一上来就用大batch size。因为图像是显微切片细节多batch太小估计噪声大batch太大则GPU显存不够而且BN层的统计量容易抖动。通常batch size取32到64。优化器用AdamW或SGD带动量预训练模型微调用AdamW更省心SGD更容易配出高精度但调参费时。学习率方面常见做法是先用0.1倍于新训练层的学习率来更新主干新加的FC层使用完整学习率。实现上可以有差异学习率也可以简单地把主干层的学习率乘0.1。我的经验是ResNet50 AdamW 初始lr 1e-4训练20个epoch前5个epoch做线性warmup然后cosine decay到1e-5效果很稳。如果你的数据增强比较强比如随机旋转、缩放、颜色抖动、cutmix那么训练周期要适当拉长。40个epoch起步比较正常。混合精度训练能用就用半精度可以让batch翻倍训练速度提升明显。不过第一次跑多标签时我建议先关AMP把整个流程跑通因为AMP可能带来数值回退排查起来麻烦。from torch.cuda.amp import autocast, GradScaler scaler GradScaler() for images, labels in train_loader: images images.to(device) labels labels.to(device) with autocast(): logits model(images) loss criterion(logits, labels) optimizer.zero_grad() scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()这里的autocast会自动化减少部分算子的精度GradScaler负责防止梯度下溢。如果你发现验证指标在某个epoch后突然掉下去先确认是不是AMP导致的数值问题再考虑关闭AMP重训。实际上HPA竞赛半精度训练的模型和全精度差别很小但速度差很多。4. 验证与推理KFold、TTA和模型集成怎么组合4.1 5折交叉验证如何搭配多标签逻辑单个模型跑出好结果是不够的竞赛里最终提交的都是多个模型或多个折的集成。5折交叉验证的每个折可以独立训练出一个模型然后把5个模型的预测结果平均起来得到更稳健的概率。多标签场景下平均操作的对象是28个标签各自的logit或概率不能用单标签的argmax逻辑。实现上先把每个折的模型保存下来推理时逐个加载。for fold in range(5): model create_model() model.load_state_dict(torch.load(fmodel_fold{fold}.pt)) model.eval() fold_logits [] with torch.no_grad(): for batch in test_loader: logits model(batch) fold_logits.append(logits.cpu()) all_fold_logits.append(torch.cat(fold_logits))最后把五个折的logits取平均再经过Sigmoid得到概率。注意不要先各自Sigmoid再平均而是先平均logits再Sigmoid这样能保留更尖锐的边界信息。如果模型之间差异很大也可以对logits做rank averaging但通常是简单平均就够。交叉验证还有一层用途是确定阈值。多标签分类你最后要对每个标签设定一个阈值判断是否存在。阈值可以从验证集上学出来对每个标签分别尝试0.3、0.4、0.5等找F1最大的点。不要统一用0.5因为稀有标签的概率普遍偏低统一阈值会把它们全切到False。4.2 TTA增强让推理时的图像匹配训练分布TTATest Time Augmentation就是把推理时的同一张图做多个几何变换分别预测后把结果合成。在HPA图像分类里TTA的收益比在普通CIFAR分类更明显因为细胞图像有旋转对称性染色位置不会因为翻转而改变语义。常用变换组合是水平翻转和垂直翻转加上原图做成三倍推理。也有人用90度旋转组合形成8倍TTA但推理时间翻倍收益却很小。实现时可以和训练时数据增强保持一致避免训练用旋转、推理不用的失配问题。我常用的TTA代码片段import torch import torchvision.transforms as T tta_transforms [ T.Compose([T.Resize((256, 256)), T.ToTensor()]), T.Compose([T.Resize((256, 256)), T.RandomHorizontalFlip(p1.0), T.ToTensor()]), T.Compose([T.Resize((256, 256)), T.RandomVerticalFlip(p1.0), T.ToTensor()]), ] def predict_with_tta(model, img, device): logits [] for tf in tta_transforms: x tf(img).unsqueeze(0).to(device) with torch.no_grad(): lg model(x) logits.append(lg) return torch.mean(torch.stack(logits), dim0)这里的RandomHorizontalFlip(p1.0)就是确定的翻转不是随机。三个变换的均值作为最终预测。注意如果原始图像是通过灰度三通道读取的TTA里不要加颜色抖动否则推理时的颜色扰动会破坏生物学信号。TTA不适合所有场景。如果你的模型已经在无翻转增强下训练得很好强加TTA反而会因为分布偏移让预测变得平滑而降低指标。一般建议在验证集上先看有无TTA的F1差别有增益再用。HPA上三倍TTA通常能带来单模F1提升0.002左右集成后提升空间小一些。4.3 加权集成与阈值寻优不同模型的性能不一样简单平均不是最优。最常见的是根据验证集F1给模型分配权重比如A模型F1 0.62B模型F1 0.58那么A权重可以调高一些。但权重差距不宜过大否则退化成只使用一个模型。另一个更稳的方案是直接对logits做几何平均即对每个标签概率取对数再平均然后再Sigmoid这样能惩罚那些极度自信的错误预测。阈值寻优的时候要注意多标签的评价指标是宏F1还是微F1。HPA竞赛使用两个指标宏观F1和微观F1最终按某种公式合并。我在调阈值时用的是自己定义的宏观F1优先因为低频标签对宏观F1影响大而高低频标签的预测概率分布差异很大。如果你只用一个全局阈值大概率是高频标签表现好低频标签被坑死。实践上我会在验证集上对每个标签独立搜索阈值。from sklearn.metrics import f1_score best_thresholds [] for i in range(num_classes): best_score 0 best_thr 0.5 for thr in np.arange(0.2, 0.8, 0.05): y_pred (valid_prob[:, i] thr).astype(int) score f1_score(valid_label[:, i], y_pred) if score best_score: best_score score best_thr thr best_thresholds.append(best_thr)这步看起来像“玄学”但其实每个标签的阈值都有意义如“细胞核”信号强阈值可以高一点稀有标签信号弱阈值要降低。搜索完毕后把阈值保存为数组推理时按标签维度应用。如果你发现某个标签的F1怎么调都上不去大概率是训练数据里这个标签的样本太少或者这个标签的染色形态和别的标签高度重叠。5. 避坑清单HPA数据上最常见的翻车现场5.1 现象验证集分数很乐观公榜一落千丈原因数据划分没有做多标签分层验证集里恰好包含了大量高频标签模型只学会了高频标签的预测或者验证集和训练集来自不同批次比如HPA数据集里同一plate的图像存在batch效应随机划分会让验证集泄漏这种batch信息。解决改用MultilabelStratifiedKFold并按样本id分组确保同一组图像不跨折。HPA数据里同一个细胞标本可能拍摄了多个视角这些视角高度相似如果不同折里出现同一标本的图像会把验证分数抬得很高公榜必翻车。所以划分时要以标本/样本组为单位而不是以单张图像为单位。5.2 现象模型把每个标签都预测成“不存在”原因多标签训练时正负样本比例差距太大模型很快发现“全预测0”也能把损失压得很低尤其是不使用pos_weight的情况下。另一个原因是损失被reductionsum放大梯度对稀有正样本不敏感。解决给BCEWithLogitsLoss设置pos_weight并且在训练日志里额外打印每个标签的正样本平均预测值。如果某类一直输出小于0.1的概率说明它被负样本淹没了。此时可以单独提高pos_weight指数或者对这类标签做随机过采样把包含稀有标签的图像重复送入训练。5.3 现象图像缩放后细胞结构细节全丢原因原图512x512细胞直径可能只有几十像素直接缩放成224x224会丢失大片细节。很多自然图像的常规操作——如最短边缩放到256再中心裁剪——在显微图像里不合适因为细胞的边缘和内部纹理都是关键特征中心裁剪会砍掉部分细胞体。解决用resize到384x384而不是224或者使用随机缩放加裁剪的组合。我在训练时用RandomResizedCropscale在0.51.0之间配合最终resize到320这样既保证目标大小变化又保留足够分辨率。推理时保持输入分辨率和训练一致否则验证和推理的尺度不对齐TTA也救不回来。5.4 现象多标签AUC不涨单标签反而涨原因损失函数和指标之间不匹配。比如你用了多标签软间隔损失但最终评测是各类别F1的平均或者模型在训练时逐步把所有标签概率拉低单标签AUC可能有提升但多标签整体的Hamming loss变差了。解决训练监控不要只看损失要每个epoch结束后跑一遍完整验证集的多标签F1和各类别AUC平均。如果发现某个类别概率持续偏低就要怀疑pos_weight是否设置正确。另外不要在训练早期用手动调整的阈值去评测早期概率分布还没稳定阈值应放在训练完成后再统一搜索。5.5 现象推理时被显存卡死原因集成多个模型时把每个模型的输入和输出都保存在GPU上显存迅速耗尽。尤其是TTA 8倍5折集成数据量瞬间放大。解决推理时不要一次性加载整个模型列表。用with torch.no_grad()包住推理并且逐个模型推理后把logits转移到CPU存储释放GPU显存。如果还卡就按batch循环推理每处理完一个batch就清空变量。另外TTA的数量可以先从3倍开始评估一下增益再决定是否增加不要盲目上8倍。6. 最后一层进阶伪标签、阈值与推理速度的取舍做到这一步你已经能从训练集得到靠谱的模型并且在验证集上找到每个标签的最优阈值。接下来想再往上提分常见动作是使用伪标签pseudo labeling用当前模型对无标签的测试集做预测把置信度很高的样本连同预测标签加入训练集重新训练一遍。HPA竞赛里伪标签有奇效因为测试集和训练集来自同分布那些模型确信度超过0.95的预测很大概率是对的。使用时要小心只能把置信度高的样本加进去而且伪标签样本权重可以降低设为原损失权重的0.5左右避免模型被自己的错误强化。另一个细腻的技巧是“标签阈值与训练目标联动”。如果最终评价F1而你只把BCE损失当训练目标训练和评价会存在一定割裂。可以尝试在训练后期把logits经过可导的sigmoid后计算一个近似F1的损失比如用soft F1做一个短时间的微调。这样模型会倾向于输出更符合F1目标的双峰分布推理时阈值搜索会更稳定。不过这个技巧对学习率很敏感一般只在最后5个epoch用1e-5的学习率做。推理速度的取舍也要提前算。Kaggle竞赛有提交时限HPA测试集有近5万张图像如果5折集成加8倍TTA每张跑完可能要好几个小时远超个人电脑的承受范围。我的习惯是用3倍TTA加5折模型单卡V100上大约可以接受。如果资源有限可以砍掉TTA改为只做水平翻转收益损失很小。更极端时把集成模型数量从5个减到3个用平均logits代替权重集成速度能提升40%而F1只掉0.003以内。我做类似竞赛有个习惯每次修改预处理或调参先在单一小模型上训练5个epoch看验证F1变化趋势确认有效再全量跑。这样能避免一次跑几小时最后发现是数据加载出了bug。还有一件事模型保存时一定带上验证指标和阈值否则你重新加载后找不到当初用什么阈值提交还要重新调一遍白白浪费算力。希望这些经验和教训能帮到你让你在显微镜图像分类这条路上少走几个来回。本文还有配套的精品资源点击获取
网站建设高端定制企业官网