ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

FixMatch半监督学习:伪标签与一致性正则化的PyTorch实现

FixMatch半监督学习:伪标签与一致性正则化的PyTorch实现 搞深度学习的同学十有八九都听过 FixMatch 这个名字。只要是做半监督学习它基本就是绕不开的基线模型。NeurIPS 2020 提出之后这个算法在 CIFAR-10、SVHN、ImageNet 这类数据集上用极少量标注数据比如每类 4 张、25 张、400 张刷出了非常接近全监督的结果。更难得的是它的实现逻辑并不复杂核心思想可以用一句话概括用弱增强样本生成高置信度的伪标签再让模型对强增强样本的预测向这个伪标签对齐。这句话拆开来看既有伪标签Pseudo-Label的影子又有一致性正则化Consistency Regularization的内核两个方法结合到一起产生了 112 的效果。这篇博文的主要内容就是带着大家把 FixMatch 从原理到 PyTorch 代码逐行过一遍。我默认你已经有 PyTorch 基础知道 DataLoader、Module、反向传播这些基本概念。如果你正在复现 FixMatch、写半监督方向的代码或者打算在自己的任务里用上这种少量标注 大量无标注的训练方式这篇文章应该能省掉你不少绕路的时间。文中涉及的代码都是我实际跑过的写法注释和解释我会放在最关键的位置让你不光能跑通还能明白每一步为什么这么写。1. 从公式看懂 FixMatch一条 loss 串起整个算法先把最核心的训练目标摆出来。FixMatch 的总损失由有监督损失和无监督损失两部分组成[ L L_s \lambda_u L_u ]这个式子看着简单但里面的每一个符号都值得细抠。(L_s) 是在有标注数据上的标准交叉熵损失(L_u) 是在无标注数据上的伪标签交叉熵损失(\lambda_u) 是控制无监督损失权重的超参数。整个算法的关键设计都在 (L_u) 的计算方式上。1.1 有监督分支和其他分类任务没有区别有监督这部分确实就是常规操作。模型接收一张有标注图片 (x_l)输出 logits然后和真实标签 (y_l) 计算交叉熵[ L_s \frac{1}{B_l} \sum CE(p(y | \text{weak}(x_l)), y_l) ]这里的 (\text{weak}(x_l)) 表示对标注图片做弱增强比如随机水平翻转和随机平移。你可能会有个疑问为什么有监督数据也要做弱增强因为 FixMatch 给标注数据和无标注数据用的是同一条数据增强流水线弱增强是为了和有标注分支的输入分布对齐。而且实际实验里这种轻度的增强对防止过拟合也有帮助尤其是标注数据量很少的时候不做增强的话模型很快就把训练集背下来了。1.2 无监督分支弱增强生成目标强增强提供学习信号无监督分支是整个算法的灵魂。对一张无标注图片 (u)我们同时做两个操作弱增强 (\alpha(u))得到的图片送入模型得到预测 (p_w p(y | \alpha(u)))。取它的最大概率对应的类别作为伪标签最大概率值作为置信度强增强 (A(u))得到的图片也送入模型得到预测 (p_s p(y | A(u)))。随后计算 (p_s) 与伪标签之间的交叉熵。但这里有个关键过滤条件只有弱增强预测的置信度超过阈值 (\tau) 时这个样本的伪标签才被采用它的无监督损失才会计入总损失。用公式表达[ L_u \frac{1}{B_u} \sum \mathbb{1}(\max(p_w) \ge \tau) \cdot CE(p_s, \hat{y}_w) ](\hat{y}_w) 是弱增强预测的 argmax 结果(\mathbb{1}(\cdot)) 是指示函数条件满足时为 1否则为 0。这个设计是 FixMatch 区别于其他半监督方法的核心。1.3 一致性 伪标签为什么能凑效在 FixMatch 出现之前半监督学习主要有两条技术路线一条是伪标签Pseudo-Labeling简单粗暴地让模型对无标注数据输出一个预测作为软标签或硬标签来训练另一条是一致性正则化代表方法有 Mean Teacher、Pi-Model 等思路是让同一个样本在不同扰动下产生相似的预测。这两条路各有问题。伪标签最大的痛点是确认偏差confirmation bias模型一旦给无标注数据打了错的标签这个错误还会在训练中自我强化一致性正则化的问题在于它只要求预测一致但模型完全可能一致地错下去缺少一个明确的学习目标。FixMatch 把两者拼在一起一致性正则化负责提供稳定的学习信号而高阈值伪标签负责提供目标值。因为阈值设得很高通常是 0.95伪标签的准确率非常高噪声被压到很低确认偏差的问题就得到了有效缓解。为了帮你把直觉建立起来我打个比方。弱增强和强增强看起来像同一个内容抄两遍本质上是让模型学到一种鲁棒的映射关系喂给它一张模糊的、被裁剪的、被遮挡的图它也应该能认出和清晰版本一致的内容。高阈值则像是一位严格批改作业的老师只有模型对弱增强版本足够有把握这条作业才被允许进入训练没把握的样本直接丢弃宁缺毋滥。这就是为什么很多实验里 FixMatch 在 CIFAR-10 上每类只有 4 张标注图时仍然能保持 90% 上下的准确率。2. 数据准备与增强流水线给模型做对照实验数据部分是复现 FixMatch 时最容易踩坑的地方。一个常见的问题是很多人拿 torchvision 默认的 CIFAR-10 训练集又当标注集又当无标注集预处理逻辑混乱最后模型根本训不动。这里我把数据组织方式和增强流水线分开讲清楚。2.1 半监督数据集怎么组织标注集和无标注集不能混首先明确两个数据集的定义标注集labeled set和无标注集unlabeled set。它们都来自同一个原始数据集区别只在于是否有可用标签。比如 CIFAR-10 原本有 50000 张训练图如果做每类 4 张的设置就从中每类挑出 4 张一共 40 张作为有标注数据其余 49960 张全部视为无标注数据它们的标签在训练时不可见。这里的关键点是无标注数据并非没有真值而是在训练过程中我们不使用它。按这个思路代码里通常这样组织# 假设原始训练集叫 train_dataset有 50000 张图 num_labeled 40 # 每类 4 张共 10 类 labels np.array(train_dataset.targets) # 按类别各取 4 张作为标注集 labeled_idx [] for c in range(10): idx_c np.where(labels c)[0] labeled_idx.extend(np.random.choice(idx_c, 4, replaceFalse)) labeled_idx np.array(labeled_idx) unlabeled_idx np.setdiff1d(np.arange(len(train_dataset)), labeled_idx) labeled_dataset Subset(train_dataset, labeled_idx) unlabeled_dataset Subset(train_dataset, unlabeled_idx)给无标注数据集构造 DataLoader 时千万注意不要设置shuffleTrue的默认行为意外重要因为无标注数据不参与有监督 loss它的顺序只影响每次采样到的未标注样本组合。常见做法是给两个数据集成对抽样保证每个训练步里标注数据和无标注数据都能被取到labeled_loader DataLoader(labeled_dataset, batch_size64, shuffleTrue, num_workers4) unlabeled_loader DataLoader(unlabeled_dataset, batch_size64, shuffleTrue, num_workers4) # 训练循环中 for (x_l, y_l), (x_u, _) in zip(labeled_loader, unlabeled_loader): pass使用zip时要特别注意两个 DataLoader 的样本数量差别很大一个只有 40 张一个有 49960 张如果直接 zip 完再结束会先走到短的那个 loader 的尽头。因此训练循环一般以步数为单位每步都从两个 loader 里各取一个 batch直到走完设定的总步数而不是以 epoch 为单位。这也是复现 FixMatch 时最容易犯的错之一。2.2 弱增强和强增强一个求稳一个求狠FixMatch 的增强设计是整个算法的标尺。弱增强解决学习目标从哪来的问题强增强解决学习信号如何逼着模型变得更鲁棒的问题。弱增强在原始论文里用的是标准翻转平移组合随机水平翻转概率 0.5配合 4 像素的随机裁剪填充。CIFAR 图片是 32x32在做平移前通常先 padding 到 40x40再随机裁剪回 32x32。PyTorch 里的写法class WeakAugment: def __call__(self, x): x transforms.RandomHorizontalFlip(p0.5)(x) x transforms.RandomCrop(size32, padding4)(x) return x强增强则要狠一些由 RandAugment 和 Cutout 两部分组成。RandAugment 会从一组图像变换比如旋转、剪切、颜色抖动、对比度调整、平移、缩放等中随机挑 N 个变换每个变换的幅度由参数 M 控制。之后再做 Cutout也就是在图像上随机挖掉一块正方形区域强度一般设为 16x16 的 patch。整体代码通常如下class StrongAugment: def __init__(self, n2, m10): self.rand_augment RandAugment(num_opsn, magnitudem) self.cutout Cutout(n_holes1, length16) def __call__(self, x): x self.rand_augment(x) x self.cutout(x) return x关于 RandAugmentPyTorch 自 1.10 版本起在torchvision.transforms.autoaugment里提供了RandAugment类直接调用即可。老版本则需要自行实现GitHub 上有不少开源的实现代码复制回来时要注意变换列表是否包含几何变换和颜色变换二者缺一不可。颜色变换能增加模型对光照的鲁棒性几何变换则迫使模型学习空间不变性Cutout 本质上是一种遮挡正则化逼着模型不全靠局部特征做判断。强弱增强的对比本质上就是在给模型做对照实验弱增强版本代表稳定的答案强增强版本代表困难的问题。模型必须学会从困难版本里提取出和稳定答案一致的语义信息。这也是半监督学习里一致性二字最直观的表达。2.3 无标注数据和标注数据的 batch 比例原始 FixMatch 论文里有个关键超参数 (\mu)表示无标注数据 batch size 和有标注数据 batch size 的比值。论文里最常用的是 (\mu 7)即有标注 batch 为 64 时无标注 batch 为 448。这样每个训练步里无标注数据贡献的信息量远大于有标注数据。毕竟在真实场景中无标注数据量大且获取成本低理应被充分利用。不过如果你显存不够或者只是跑通流程验证代码可先从 (\mu 1) 开始后面再调大。实际操作中把弱增强和强增强看成两个独立的分支分别生成两个无标注 batchweak_x_u weak_augment(x_u) strong_x_u strong_augment(x_u)同一个x_u经过两条不同的增强路径得到两个版本的输入。这里的顺序是先做弱增强再做强增强还是分别独立做对结果影响不大但务必保证弱增强和强增强不在同一个对象上原地修改否则会出现两个版本互相污染的问题。3. 模型定义与训练骨架选 backbone 和优化器的经验FixMatch 是一个模型无关的训练框架理论上任何分类 backbone 都可以接。但既然要复现论文效果了解原论文的配置和背后的原理还是有必要的。3.1 backbone 选择WideResNet 还是简化 CNN原论文中选用的 backbone 是 WideResNetWRN-28-2 或 WRN-28-8这是半监督和自监督领域的老熟脸了。它相比 ResNet 的主要改动是把卷积层的宽度扩大即通道数乘以一个缩放因子同时加深每个 block 里的卷积层数在参数量可控的情况下显著提升特征表达力。WideResNet 结构里也常带有 DropBlock 或 dropout能更好地防止模型在小标注集上过拟合。但如果你想快速验证 FixMatch 代码逻辑完全可以用一个轻量 ResNet-18 或者更简单的 CNN 代替。原因很简单FixMatch 的核心创新不在网络结构而在训练目标的设计。先在小网络上跑通完整流程再切换到 WideResNet 做最终实验是效率最高的路径。我自己复现的时候第一版代码就用了一个 5 层卷积的简易网络训练 200 个 epoch 后照样能看到无监督损失带来的精度提升。模型输出层的维度必须和类别数一致这是废话但值得强调。伪标签的生成依赖模型在某个类上的置信度超过阈值如果类别数搞错后面代码会直接报 shape 不匹配的错排查起来相当浪费时间。3.2 EMA 要不要用怎么用指数移动平均Exponential Moving AverageEMA是很多半监督方法里的标配组件。它维护一份模型参数的滑动平均副本用这个副本去做推理或衡量模型状态通常比直接用训练中的参数要稳定很多。FixMatch 原始论文本身没有把 EMA 当作核心卖点但我在实验里发现加上 EMA 后验证集的精度波动会减小不少。原因很容易理解小标注数据下模型参数在训练后期仍然会大幅震荡滑动平均能把这些震荡平滑掉让网络输出更稳定。实现 EMA 并不复杂核心就是维护一个影子模型ema_model copy.deepcopy(model) def update_ema(ema_model, model, decay0.999): with torch.no_grad(): for ema_param, param in zip(ema_model.parameters(), model.parameters()): ema_param.data.mul_(decay).add_(param.data, alpha1 - decay)每次 optimizer.step() 之后调用一次update_ema。decay 通常取值 0.999 或 0.9999。注意两件事第一EMA 模型不需要计算梯度所以它参与前向推理时务必用torch.no_grad()包起来第二如果模型里用了 BatchNormEMA 模型的 BatchNorm 统计量并不会随着参数滑动自动更新这时应该在每个 epoch 结束时用训练数据跑一次前向把 BN 的 running_mean 和 running_var 更新到 EMA 模型上否则 EMA 模型的推理结果可能不准。3.3 优化器与学习率调度的配置经验原论文用的是 SGD 优化器momentum 为 0.9weight decay 为 5e-4初始学习率 0.03配合 cosine 学习率衰减。PyTorch 里的标准写法optimizer torch.optim.SGD( model.parameters(), lr0.03, momentum0.9, weight_decay5e-4, nesterovTrue ) scheduler torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_maxtotal_steps, eta_min0.0 )关于这个配置有几点经验可以分享。第一SGD cosine 在 FixMatch 上的表现通常优于 Adam 一族。半监督训练任务中学习率的动态范围变化很大cosine 衰减能帮模型在后期稳定收敛Adam 则往往因为二阶动量累积导致后期调整能力减弱。但这不绝对如果你的任务本身比较特殊用 AdamW 也完全可以只是需要重新调参。第二total_steps不是 epoch 数而是总的训练步数。比如总训练步数是 2^20大约 104 万步这个量级在 CIFAR-10 上大约对应 1000 多个 epoch。实际复现时很少有人真的跑满 100 万步跑个 20~40 万步已经能看到明显效果。第三如果标注数据实在太少可以考虑在前若干个 epoch 用 warmup把学习率从 0 线性升到 0.03但我实测下来 FixMatch 对 warmup 的依赖不大可加可不加。4. 核心训练循环代码逐行拆解伪标签、mask 与 loss进入正题。这里我会给出一个完整的训练循环代码片段并逐段解释逻辑。为了让代码尽量通用我假设模型输入是x输出是 logitsce_loss是标准的交叉熵函数。4.1 生成伪标签并做阈值过滤对一个无标注 batch先把弱增强版本的图片送入模型得到 logits然后执行三个操作用 softmax 把 logits 转成概率分布用torch.max同时取出最大概率值max_probs和对应的类别索引pseudo_label后者就是伪标签用最大概率值和阈值 (\tau) 比较生成一个布尔 mask。代码片段如下with torch.no_grad(): logits_weak model(weak_x) probs_weak torch.softmax(logits_weak, dim1) max_probs, pseudo_label torch.max(probs_weak, dim1) mask (max_probs tau).float()这里最容易被忽略的是torch.no_grad()。伪标签生成过程只是用模型当前的预测来给训练数据打标这一步不需要反向传播因此必须锁住梯度。如果不加no_grad()模型对弱增强样本的前向计算也会被纳入计算图既浪费显存又可能在反向传播时导致梯度流到不该去的地方。tau通常取 0.95。这个值不是拍脑袋定的论文作者做过大量实验0.95 在 CIFAR-10 和 SVHN 上表现最好。阈值越高伪标签的准确率越高但有资格参与训练的样本就越少阈值过低伪标签噪声变大模型容易学到错误信息。这个平衡是半监督学习中利用 vs 探索问题的具体体现。4.2 mask 如何参与 loss 计算有了 mask接下来计算无监督损失。对强增强版本样本模型输出 logits再求和伪标签的交叉熵但在最后要乘上 mask。这里有一个实现上的细节差异值得专门讲一下。如果你直接使用 PyTorch 的F.cross_entropy(logits, pseudo_label, reductionmean)它会把整个 batch 的 loss 平均。但我们只想让 mask 为 1 的样本贡献 loss所以必须改成reductionnone先得到逐样本 loss再手动乘上 mask 求平均loss_unsup F.cross_entropy( logits_strong, pseudo_label, reductionnone ) loss_unsup (loss_unsup * mask).mean()这一步是整个 FixMatch 实现的胜负手。如果你犯了低级错误直接把mask和 loss 相乘的顺序写反或者在reductionnone之前就缩成了标量那么 mask 就完全失效了。另外还要注意mask的类型是 float这样乘法运算才能进行。如果你想让代码更保险也可以对mask.sum()做除法而不是mask.mean()但如果某个 batch 的 mask 全为 0就会出现除以 0 的问题。用.mean()的写法在数学意义上等价于对被采纳样本求平均只是样本数为 0 时结果为 0不会崩。4.3 完整训练循环把有监督分支和无监督分支拼到一起就得到了整个 FixMatch 训练循环的主干model.train() for step in range(total_steps): # 取出有标注 batch 和无标注 batch try: (x_l, y_l) next(labeled_iter) except StopIteration: labeled_iter iter(labeled_loader) (x_l, y_l) next(labeled_iter) try: (x_u, _) next(unlabeled_iter) except StopIteration: unlabeled_iter iter(unlabeled_loader) (x_u, _) next(unlabeled_iter) x_l, y_l x_l.to(device), y_l.to(device) x_u x_u.to(device) # 强弱增强 weak_x_u weak_augment(x_u) strong_x_u strong_augment(x_u) # 有监督 loss logits_l model(x_l) loss_sup F.cross_entropy(logits_l, y_l) # 无监督伪标签生成 with torch.no_grad(): logits_weak model(weak_x_u) probs_weak torch.softmax(logits_weak, dim1) max_probs, pseudo_label torch.max(probs_weak, dim1) mask (max_probs tau).float() # 无监督 loss logits_strong model(strong_x_u) loss_unsup F.cross_entropy( logits_strong, pseudo_label, reductionnone ) loss_unsup (loss_unsup * mask).mean() loss loss_sup lambda_u * loss_unsup optimizer.zero_grad() loss.backward() optimizer.step() scheduler.step() update_ema(ema_model, model)这段代码的每个点几乎都有讲究我依次说明。第一标注和无标注两个 DataLoader 各自维护独立的迭代器每当一个 epoch 跑完就重新初始化。这样即使两者长度差很多训练步数也能精确控制不受 epoch 对齐的限制。比前面直接用zip更稳。第二无标注数据的原始标签_直接丢弃。代码里看起来是没用到标签但实际数据集的标签还保存在内存里只是我们没有使用。如果你把无标注数据集的标签传入模型参与任何计算那就变成有监督了结果会虚高毫无意义。第三lambda_u是超参数论文里取 1.0。有监督 loss 和无监督 loss 直接相加意味着模型在训练步中同时收到两股梯度信号。按照真实经验如果你发现某个 batch 里 mask 全为 0无监督 loss 为 0 是正常现象不代表代码有 bug这在训练初期尤其常见模型还不够自信时会频繁发生。第四update_ema(ema_model, model)放在每次 optimizer.step() 之后这是标准做法能保证 EMA 模型始终是历史参数的平滑版本而不是当前参数的前一版。4.4 模型评估阶段别忘了 mask训练结束后评估模型时一般直接使用 EMA 模型。推理逻辑很常规不需要任何伪标签和 mask 相关的代码ema_model.eval() correct, total 0, 0 with torch.no_grad(): for images, labels in test_loader: images, labels images.to(device), labels.to(device) outputs ema_model(images) _, predicted torch.max(outputs, dim1) total labels.size(0) correct (predicted labels).sum().item() print(fAccuracy: {correct / total:.4f})但千万记得在评估前把ema_model切到eval()模式否则 BatchNorm 和 Dropout 的行为还是训练模式结果会明显变差。我见过不少人把训练和评估写在一个函数里切换模式时漏掉了 EMA 模型导致验证集准确率一会高一会低排查很久才发现是这种低级错误。5. 踩坑记录与参数调优建议来自复现现场的经验跑通代码只是开始把精度调上去才是真正花时间的地方。以下是我复现过程中实际遇到的典型问题给出排查思路和推荐解决措施。5.1 训练初期 loss 为 0 是正常的但一直为 0 就不正常很多第一次跑 FixMatch 的同学看到前几个 step 里无监督 loss 一直是 0.0会怀疑是不是 mask 有问题。训练初期模型权重接近随机初始化对弱增强样本的预测置信度普遍低于 0.95mask 全为 False 是完全正常的。一般跑过几百步之后模型学到基础特征部分无标注样本的置信度开始超过阈值无监督 loss 才会慢慢变大。但如果跑了几千步无监督 loss 仍然恒为 0就需要检查几件事确认mask (max_probs tau).float()里tau是不是写成了常数 0.95而不是一个被你错误赋值成 1.0 的变量确认模型输出维度是否为类别数比如 CIFAR-10 是 10 维确认弱增强图片和强增强图片没有搞反如果弱增强实际使用的是强增强伪标签会难以越过 0.95 的阈值。还有一种情况也要警惕预训练模型或非随机初始化模型可能一开始置信度就很高导致 mask 几乎是全 1这时的伪标签质量其实很差模型会迅速被噪声带偏。如果使用预训练 backbone应该先把阈值调低或者在无监督 loss 上加上温启动warmup策略。5.2 阈值 (\tau)、权重 (\lambda_u) 和 batch 比例 (\mu) 怎么搭配这三个超参数会互相影响不建议同时猛调。先从论文默认值开始(\tau0.95)、(\lambda_u1.0)、(\mu7)然后每次只动一个量。先看 (\tau)。如果模型验证集准确率一直上不去但伪标签的 mask 覆盖率和预期差不多可以尝试把 (\tau) 降到 0.9观察训练初期无监督 loss 是否更快出现。调低 (\tau) 相当于放宽了谁有资格提供学习信号的门槛让更多无标注样本进入训练但也引入了更多噪声。我的经验是数据集类别越多、图片越复杂阈值应该越低因为复杂图片的类别边界天然模糊强制保持 0.95 会让大量样本被过滤掉无监督 loss 形同虚设。CIFAR-100 这类 100 类任务上0.8~0.9 的阈值有时表现更好。再看 (\lambda_u)。它控制无监督损失的相对强度。(\lambda_u) 太大模型会过度依赖无标注信号一旦伪标签里混入少量错误错误就会被放大太小则无监督分支形同虚设。论文默认 1.0 是个均衡点。如果你的标注数据特别少可以试试把 (\lambda_u) 设为 2 或 3让无标注数据的信息更充分地被利用。最后是 (\mu)。显存允许的情况下把无标注 batch 调大通常能加速收敛。但注意在 PyTorch 里batch 越大每个 step 的 GPU 显存占用越高如果你只有单张 1080Ti(\mu7) 配 64 的标注 batch 可能会爆显存降到 (\mu2) 或 (\mu4) 即可。5.3 伪标签噪声太大怎么解决确认偏差的预防即使有 0.95 的阈值训练后期模型偶尔也会给无标注样本打错标签错误会累积这就是确认偏差。FixMatch 已经比纯伪标签方法好很多但并不能完全消除这个问题。下面几个手段可以在不改动算法核心的前提下有效缓解在标签噪声较大的初期对无监督 loss 做 ramp-up。比如前 2000 步让 (\lambda_u) 从 0 线性升到目标值给模型一个先用有监督信号学好基础的时间窗口给未标注数据加上分布约束比如统计模型在验证集上的类别预测分布强制无标注分支输出的类别分布不要偏离太大这属于偏分发正则化的思路实现上会复杂一些定期在验证集上评估如果发现验证集准确率连续多个 epoch 下降可以先停止训练用当前最好的模型状态重新生成伪标签再继续训练。坦白讲这些手段对最终精度提升的作用都不是决定性的但它们能让训练过程更稳定各种随机种子下跑出来的结果不会像过山车一样忽高忽低。5.4 常见问题速查表为了让你排查问题时更高效我整理了一张速查表覆盖我在复现过程中遇到的和朋友遇到过的常见情况。现象可能原因解决建议无监督 loss 一直为 0阈值太高或模型太弱检查 tau 是否为 0.95降低 tau 或增加 warmuploss 为 NaN学习率过大把 lr 降到 0.01 或更低排查训练精度高但验证精度低标注数据过少引发过拟合启用更强的数据增强或者调大 weight decay验证精度波动剧烈单 batch 有监督信号太弱使用 EMA 模型评估尝试调大 batch size训练速度太慢大量无标注样本被 mask 过滤提高弱增强的强度或适当降低阈值显存不足无标注 batch 过大降低 mu 值或使用梯度累积模拟更大 batch两个 loader 长度不匹配总步数失控使用迭代器方式控制训练步数这张表不能覆盖所有问题但如果你遇到的和表中某一行吻合按对应的建议去试错大概率能走通。5.5 关于分布式训练和多卡的一点补充如果你的任务规模上了 ImageNet 级别单卡训练会非常慢这时需要用DistributedDataParallelDDP把模型分布到多张卡上。FixMatch 在 DDP 下的实现本身没有特殊之处但有几个细节需要留意。伪标签的生成在每张卡上独立进行这没有问题但不同卡上的 BatchNorm 统计量在分布式模式下默认是不同步的。对于半监督训练这会导致每张卡上模型看到的数据分布不完全一致影响无监督 loss 的稳定性。建议开启torch.nn.SyncBatchNorm.convert_sync_batchnorm(model)做同步 BatchNorm。这个操作会带来额外的通信开销但收益在标注数据很少的场景下非常明显。另一个细节是数据增强的随机种子。多卡训练时如果同一张无标注图片在两张卡上经过了不同的增强其对应的伪标签和强增强版本可能属于不同的数据流逻辑上依然说得通但为了复现方便建议把每个 step 的随机种子固定下来或者在 DataLoader 的worker_init_fn里为每个 worker 设置独立种子保证每次运行结果可重复。写在最后的一些体会FixMatch 的代码实现细节并不复杂但它能在半监督学习领域占据重要位置靠的是对各种设计细节的精准把握弱增强提供目标、强增强提供信号、高阈值过滤噪声这三个设计环环相扣。我自己在不同数据集上跑过的感受是想要完全复现论文里的精度数字除了把代码写对更重要的是对超参数保持耐心多试几组组合对比。单独看每个超参数的影响都不大合在一起却可能带来好几个百分点的差异。如果你正在做半监督方向的研究或者手头有一个标注数据稀缺的任务FixMatch 绝对是一个值得花时间去吃透的基线。按这篇文章里的代码和调参思路走下来你应该能在自己的任务上快速跑出一个可靠的结果后续再在这个基础上改进也会顺手很多。
返回列表