ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

FixMatch半监督学习算法详解与PyTorch代码实现

FixMatch半监督学习算法详解与PyTorch代码实现 FixMatchNeurIPS 2020作者 Sohn 等人是半监督学习领域一篇被大家反复聊的论文。算法本身并不复杂一句话概括就是用弱增强版本的预测为同一张无标签图片的强增强版本生成伪标签再用一个固定阈值筛出高质量伪标签让模型去学习这些伪标签。思路听着很直白但真正拿 PyTorch 从头把整套代码复现一遍中间能踩的坑一点不少。我自己在复现时就被数据增强的执行顺序、伪标签 mask 的处理、无监督 loss 的除法细节这些问题绊倒过好几次。这篇文章就当成一份 FixMatch 的 PyTorch 代码详解来写。我会从环境搭建、项目目录设计开始逐步拆解 Dataset、数据增强、Wide ResNet、伪标签生成、Loss 组合、训练循环最后把常见问题汇总成一个排查清单。想自己动手复现的朋友可以照着一步步跟有 PyTorch 基础但还没怎么接触过半监督学习的读者也能看懂。代码部分我会尽量给出可运行的片段并把关键参数的计算思路讲清楚。1. 先把FixMatch的核心思路吃透1.1 半监督学习要解决的痛点是什么半监督学习的典型场景是手里有一批带标签数据同时还有一大堆没有标签的数据。在医疗影像、工业质检、用户行为这种场景里标注成本通常很高而原始数据本身往往是海量的。纯用有标签数据训练模型很容易过拟合如果直接把无标签数据丢掉又太浪费。早期的一类做法是伪标签Pseudo-Labeling让模型先预测无标签样本把置信度最高的预测当成“软标签”放回去继续训练。问题在于模型犯错时会把错误当成正确样本去强化导致错误越滚越大。另一类做法是一致性正则化Consistency Regularization对同一张无标签图片做不同扰动希望模型在两个扰动版本上输出一致的概率分布。它能利用无标签数据但对“预测有多可信”没有显式判断容易在训练初期学习到噪声。FixMatch 的高明之处在于它把这两条线缝到了一起一致性正则化负责让模型在弱增强和强增强版本上保持一致伪标签机制负责筛选出高置信度的预测作为监督信号。两者互相约束既解决了伪标签方法容易自我强化的毛病又比纯一致性正则化收敛得更快、更稳。1.2 FixMatch的算法流程与代码对应官方论文里的流程可以拆成五步对无标签图片分别做弱增强和强增强得到x_uw和x_us。模型对弱增强版本x_uw做前向推理得到预测概率。取最大概率max(p)作为置信度如果max(p)大于阈值tau论文里通常取 0.95那么把对应的类别argmax(p)当作伪标签。模型对强增强版本x_us做前向推理计算与伪标签之间的交叉熵。把有标签数据的交叉熵和这个无监督交叉熵加权相加共同更新模型。这段流程映射到代码上就是后面我即将展开的 Dataset、模型、训练循环三大模块。选择用 PyTorch 实现一是因为它生态成熟CIFAR-10、SVHN 这类常用数据集都有现成接口二是因为 FixMatch 官方仓库本身也是 PyTorch 写的后续想对照官方训练配置调参时比较省力。2. 从零搭建FixMatch的PyTorch项目目录2.1 PyTorch环境配置建议建议直接用官方源创建 conda 环境Python 3.9 以上就可以。PyTorch 版本不用追求最新但尽量选稳定版2.0 之后的版本对自动混合精度、数据加载的优化都做得不错。安装命令直接去 PyTorch 官网根据 CUDA 版本复制即可这套流程已经被写烂了我这边默认读者已经装好了 CUDA 驱动和对应版本的 PyTorch。除了 PyTorch 本身还需要torchvision和tensorboard。torchvision提供数据集接口和基础数据增强工具tensorboard用于记录训练曲线。如果是用 pip 维护依赖可以写到 requirements.txt 里torch2.0.0 torchvision0.15.0 tensorboard2.12.0 numpy装环境时有一个细节容易忽略PyTorch 版本和 CUDA 版本必须和本机驱动匹配否则会出现CUDA error: no kernel image is available这种问题。如果驱动版本较老就降级到 CUDA 11.8 的 PyTorch 包不必非追新版。2.2 项目目录结构与核心模块划分FixMatch 的代码结构不需要设计得太花哨保持职责清晰最重要。我这次复现时用的是下面这套目录fixmatch/ ├── config.py # 超参数配置 ├── data/ │ ├── __init__.py │ ├── dataset.py # Dataset定义、标签划分 │ └── transforms.py # 弱增强、强增强定义 ├── models/ │ ├── __init__.py │ └── wideresnet.py # Wide ResNet模型 ├── utils/ │ ├── __init__.py │ ├── meter.py # 平均指标统计 │ └── misc.py # 随机种子、EMA等工具 ├── train.py # 训练入口把数据增强单独拆成transforms.py很重要。FixMatch 的代码里数据增强是整个算法最核心的“调味料”弱增强和强增强分别服务于不同的分支如果堆在 Dataset 里会让代码变乱。模型单独一个文件方便以后替换成 ResNet 或者其他结构。配置统一放在config.py里而不是散落在各个脚本中这样调参时只改一处。3. 数据集与数据加载半监督数据怎么构造才对3.1 Dataset类的核心写法与标签划分FixMatch 在 CIFAR-10 上最常用的设置是 4000 张有标签样本也就是每个类别 400 张其余 46000 张都当无标签数据。划分的时候必须确保每个类在标签集中都有出现不能直接随机打乱后截前 4000否则某个类别可能完全消失。实现时我先拿到 CIFAR-10 的全部训练数据按类别索引分组在每个类别内部打乱顺序然后取前 400 张作为 labeled其余作为 unlabeledimport numpy as np import torchvision def split_semi_dataset(root, num_per_class400): full_set torchvision.datasets.CIFAR10(rootroot, trainTrue, downloadTrue) labels np.array(full_set.targets) labeled_idx, unlabeled_idx [], [] rng np.random.default_rng(42) for cls in range(10): # CIFAR-10 共10类 cls_idx np.where(labels cls)[0] cls_idx rng.permutation(cls_idx) labeled_idx.extend(cls_idx[:num_per_class]) unlabeled_idx.extend(cls_idx[num_per_class:]) labeled_idx np.array(labeled_idx) unlabeled_idx np.array(unlabeled_idx) return full_set.data[labeled_idx], labels[labeled_idx], full_set.data[unlabeled_idx], labels[unlabeled_idx]unlabeled分支其实不需要真实标签但保留下来可以在测试时粗略检查无标签数据的难度分布。num_per_class400对应 CIFAR-10 的 4000 标签设置这个参数直接影响实验效果论文里有多个档位复现时按需调整。Dataset 类需要区分训练模式和测试模式。训练模式下labeled 分支返回一张弱增强图片和标签unlabeled 分支返回同一张图片的弱增强版本和强增强版本测试模式下只返回原始图片和标签class SemiDataset(Dataset): def __init__(self, images, labels, split, weak_transform, strong_transformNone): self.images images self.labels labels self.split split # labeled / unlabeled / test self.weak_transform weak_transform self.strong_transform strong_transform def __len__(self): return len(self.images) def __getitem__(self, idx): img self.images[idx] label self.labels[idx] if self.split labeled: return self.weak_transform(img), label elif self.split unlabeled: img_w self.weak_transform(img) img_s self.strong_transform(img) return img_w, img_s else: return img, label注意transforms的输入必须是PIL.Image类型所以我在split_semi_dataset中没有提前把数组转成 tensor否则后续增强会报错。这也是新手容易踩的坑之一。3.2 DataLoader如何同时处理有标签和无标签分支有标签和无标签数据量相差很大如果直接把他们放进同一个 DataLoader需要做 batch-wise 的拼接。官方实现是把两个独立 DataLoader 用zip配对这样 labeled 一个 batch、unlabeled 一个 batch互不干扰代码也干净labeled_loader DataLoader( labeled_dataset, batch_sizeargs.batch_size, shuffleTrue, num_workers4, drop_lastTrue, ) unlabeled_loader DataLoader( unlabeled_dataset, batch_sizeargs.batch_size * args.mu, shuffleTrue, num_workers4, drop_lastTrue, )mu是 unlabeled batch 和 labeled batch 的倍数论文默认设为 7。也就是说如果 labeled batch 是 64unlabeled batch 就是 448。这个倍数直接影响无监督 loss 在总 loss 里的比重训练时如果显存不够可以先降到 2 来做验证。drop_lastTrue是为了防止最后一个 batch 形状不对导致训练中断。对 unlabeled 分支尤其要小心因为它返回的是(img_w, img_s)元组如果最后一个 batch 数据量不够zip 对不齐也会出问题。实际跑起来后我用zip(labeled_loader, unlabeled_loader)完成配对训练循环会一直运行到较短的那个 DataLoader 结束。4. 模型结构用PyTorch实现Wide ResNet4.1 精简版Wide ResNet代码FixMatch 论文在 CIFAR-10 上默认使用的是 Wide ResNet-28-2也就是深度 28、宽度倍数为 2。虽然可以调用 torchvision 的 ResNet 类但 Wide ResNet 的通道结构并不完全等价最好自己写一个精简版本。我这里给出一个去掉了复杂组件的可运行版本保留 BasicBlock 和整体前向逻辑import torch import torch.nn as nn import torch.nn.functional as F def conv3x3(in_planes, out_planes, stride1): return nn.Conv2d(in_planes, out_planes, kernel_size3, stridestride, padding1, biasFalse) class BasicBlock(nn.Module): def __init__(self, in_planes, out_planes, stride1, drop_rate0.0): super().__init__() self.bn1 nn.BatchNorm2d(in_planes) self.conv1 conv3x3(in_planes, out_planes, stride) self.bn2 nn.BatchNorm2d(out_planes) self.conv2 conv3x3(out_planes, out_planes) self.drop_rate drop_rate if stride ! 1 or in_planes ! out_planes: self.shortcut nn.Sequential( nn.Conv2d(in_planes, out_planes, kernel_size1, stridestride, biasFalse), nn.BatchNorm2d(out_planes), ) else: self.shortcut nn.Identity() def forward(self, x): out F.relu(self.bn1(x)) out self.conv1(out) out F.relu(self.bn2(out)) if self.drop_rate 0: out F.dropout(out, pself.drop_rate, trainingself.training) out self.conv2(out) out self.shortcut(x) return out class WideResNet(nn.Module): def __init__(self, depth28, widen_factor2, num_classes10, drop_rate0.0): super().__init__() assert (depth - 4) % 6 0 num_blocks (depth - 4) // 6 k widen_factor self.conv1 conv3x3(3, 16) self.layer1 self._make_layer(16, 16 * k, num_blocks, stride1, drop_ratedrop_rate) self.layer2 self._make_layer(16 * k, 32 * k, num_blocks, stride2, drop_ratedrop_rate) self.layer3 self._make_layer(32 * k, 64 * k, num_blocks, stride2, drop_ratedrop_rate) self.bn nn.BatchNorm2d(64 * k) self.fc nn.Linear(64 * k, num_classes) def _make_layer(self, in_planes, out_planes, num_blocks, stride, drop_rate): layers [] layers.append(BasicBlock(in_planes, out_planes, stridestride, drop_ratedrop_rate)) for _ in range(1, num_blocks): layers.append(BasicBlock(out_planes, out_planes, stride1, drop_ratedrop_rate)) return nn.Sequential(*layers) def forward(self, x): out self.conv1(x) out self.layer1(out) out self.layer2(out) out self.layer3(out) out F.relu(self.bn(out)) out F.adaptive_avg_pool2d(out, 1) out out.view(out.size(0), -1) out self.fc(out) return out这里有两个地方容易写错。一是 shortcut 加的是x本身不是经过归一化后的结果跳跃连接一定要跨过 BN 和 ReLU二是 BasicBlock 里的 BN 和 ReLU 顺序和普通 ResNet 略有不同Wide ResNet 用的是先 BN 再 ReLU 的 pre-activation 结构不能直接照搬经典 ResNet 的代码。4.2 为什么FixMatch选择Wide ResNet直觉上 ResNet 深度更深理论上拟合能力更强但在半监督场景里问题偏向“利用大量无标签数据的分布假设”而不是单纯“拟合更多有标签样本”。Wide ResNet 的优势在于通过增加通道数来扩大表达空间同时保持适当的深度训练更稳定、收敛更快还方便做 dropout。对于 CIFAR-10 这种 32x32 的小分辨率输入过深的网络反而会因为感受野和梯度问题变得难调。FixMatch 论文在 CIFAR-10 上用的 WRN-28-2参数量不大单张 V100 完全可以跑动。如果自己复现时想加快迭代速度可以先换 WRN-16-2 跑通整条流程验证无误后再切回论文配置。模型文件里保留widen_factor这个参数切换宽度只需要改一处配置。5. 核心算法详解增强、伪标签与Loss组合5.1 弱增强和强增强的代码实现FixMatch 里弱增强和强增强的差别本质上是“轻扰动”和“重扰动”。弱增强用的是常见的随机裁剪加随机水平翻转对语义内容影响极小。强增强采用 RandAugment这个策略是从一系列图像变换里随机抽 N 个操作并按固定幅度 M 执行带来的扰动要剧烈得多。在 CIFAR-10 的 32x32 图片场景下我的实现是这样的from torchvision import transforms from torchvision.transforms import RandAugment def get_transforms(mean, std): weak transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean, std), ]) strong transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), RandAugment(num_ops2, magnitude10), transforms.ToTensor(), transforms.Normalize(mean, std), ]) return weak, strong这里需要补充说明的是RandAugment 的官方实现要求输入是 PIL 或 tensor-based 的图像在ToTensor之前调用才可以避免转换为 tensor 后再做增强的兼容性问题。num_ops2表示每张图随机选 2 个操作magnitude10对应官方代码的强度。论文里还提到使用 Cutout如果加上的话一般放在 Normalize 之后对 tensor 操作transforms.RandomErasing(p1.0, scale(0.02, 0.16), ratio(0.3, 3.3))我在复现时发现Magnitude 调得过大反而会导致训练不稳定模型往往在强增强版本上看不清目标本身。CIFAR-10 这种低分辨率场景下magnitude10 已经够强如果后续换到更复杂的数据集可以先从 5 开始调。5.2 伪标签生成与置信度阈值筛选逻辑核心的无监督分支在训练循环里是这样实现的with torch.no_grad(): logits_uw model(images_u_w) probs_uw torch.softmax(logits_uw, dim1) max_probs, pseudo_label torch.max(probs_uw, dim1) mask max_probs.ge(threshold).float() logits_us model(images_u_s) loss_u (F.cross_entropy(logits_us, pseudo_label, reductionnone) * mask).sum() loss_u loss_u / max(mask.sum(), 1.0)整个with torch.no_grad()块是不可省的。伪标签来自弱增强分支的预测它只作为监督信号不需要回传梯度。如果不包no_grad梯度会从弱增强分支一路传回模型导致梯度计算路径混乱显存占用直接翻倍训练速度和稳定性都会受影响。max_probs.ge(threshold)是筛选逻辑的核心threshold0.95意味着只有当模型对弱增强版本的预测概率达到 95% 才认为可信。这个阈值非常关键设得太低会把错误伪标签当正确信号设得太高会出现 unlabeled 分支有效样本过少、无监督 loss 趋近于 0 的问题。我在 CIFAR-10 4000 标签实验里统计过训练初期 mask 覆盖率大概只有 20%-30%随着训练推进会逐步上升到 60%-70%这个趋势是正常的。有一种实现会写成mask (max_probs threshold)这样会得到 bool 张量参与乘法时需要先转成 float否则会出现类型错误。上面实现的ge返回 bool 后立刻float()转换直接避免了后续问题。5.3 有监督Loss与无监督Loss的组合方式有监督部分就是标准分类交叉熵logits_l model(images_l) loss_s F.cross_entropy(logits_l, targets_l)总 loss 用加权系数lambda_u把两部分组合起来论文默认lambda_u1.0loss loss_s args.lambda_u * loss_u这个lambda_u和 unlabeled batch 的倍数mu不同mu影响数据层面的比例lambda_u影响梯度层面的权重。我在复现时先按论文固定 1.0实验稳定后再尝试把它提成 2 或 5 观察变化。要注意如果调整mu或lambda_u最好每次只改一个变量不然无法判断到底是数据结构变化还是权重变化带来精度差异。组合 loss 时还有一个隐藏细节reductionnone之后手动除以mask.sum()而不是直接取.mean()。原因在于无监督 batch 里有很多低置信度样本如果直接用mean这些低置信度样本会把 loss 拉低而有效伪标签带来的梯度也会被稀释。我实际跑过对比在训练初期用mean的错误实现会让模型一直等待伪标签“变多”收敛明显变慢用有效的sum / mask.sum()之后loss 数值更合理训练曲线也更平滑。6. 训练流程完整解析优化器、调度器与主循环6.1 优化器与学习率调度怎么配论文使用的优化器是 SGDmomentum 0.9weight decay 5e-4Nesterov 开启。初始学习率在 CIFAR-10 4000 labels 场景下是 0.03这个数字和 batch size 绑定如果增大 batch size学习率通常也要等比放大。学习率调度采用 cosine schedule并配合前几个 epoch 的 warmup。warmup 的作用是防止大学习率在刚开始时把模型参数“甩飞”特别是伪标签在网络还没见过足够多数据时非常不可靠。我的实现方式之一是手动控制def get_lr(epoch, base_lr, warmup_epochs, total_epochs): if epoch warmup_epochs: return base_lr * (epoch 1) / warmup_epochs progress (epoch - warmup_epochs) / max(1, total_epochs - warmup_epochs) return 0.5 * base_lr * (1 math.cos(math.pi * progress))每轮训练前从上面拿当前学习率并给优化器赋值。也可以直接用 PyTorch 的CosineAnnealingLR但 warmup 部分需要额外包一层自定义 wrapper我嫌麻烦干脆手写。关于 EMA指数移动平均官方实现会对模型参数做 EMA并用于最终测试。这个机制可以平滑参数波动在半监督场景里确实能带来一点稳定收益。实现要点是维护一个影子模型每次参数更新后按ema_decay论文用 0.999同步影子模型参数。要注意 BatchNorm 的 running mean 和 running variance 并不参与 EMA否则 BN 统计信息会失真。6.2 完整训练循环逐段分析下面是一段可以在单卡上直接运行的简化训练循环保留核心逻辑for epoch in range(epochs): train_loss_s AverageMeter() train_loss_u AverageMeter() mask_ratio AverageMeter() lr get_lr(epoch, base_lr, warmup_epochs, epochs) set_lr(optimizer, lr) model.train() for (images_l, targets_l), (images_u_w, images_u_s) in zip(labeled_loader, unlabeled_loader): images_l images_l.cuda() targets_l targets_l.cuda() images_u_w images_u_w.cuda() images_u_s images_u_s.cuda() logits_l model(images_l) loss_s F.cross_entropy(logits_l, targets_l) with torch.no_grad(): logits_u_w model(images_u_w) probs torch.softmax(logits_u_w, dim1) conf, pseudo_label torch.max(probs, dim1) mask conf.ge(threshold).float() logits_u_s model(images_u_s) loss_u (F.cross_entropy(logits_u_s, pseudo_label, reductionnone) * mask).sum() loss_u loss_u / max(mask.sum(), 1.0) loss loss_s lambda_u * loss_u optimizer.zero_grad() loss.backward() optimizer.step() train_loss_s.update(loss_s.item()) train_loss_u.update(loss_u.item()) mask_ratio.update(mask.mean().item())zip配对会以较短 DataLoader 为准因此两个 loader 的 epoch 长度必须设计成近似一致否则每个 epoch 结束时会有一批无标签数据没被利用上。在 CIFAR-10 4000 labels 设置下labeled 有 4000 条unlabeled 有 46000 条batch 分别是 64 和 64*7448所以两者每 epoch 的 step 数不同。如果严格要每个 epoch 都让全部数据过一遍需要处理剩余样本。实际复现里我把 unlabeled 的 DataLoader 设置为drop_lastTrue然后用 labeled batch 数来控制循环轮数这样稍微损失一点数据利用率但代码清晰很多。显存方面如果一张卡只跑mu1或mu2一般 12GB 显存够用。要完整跑mu7建议 24GB 显存或者用梯度累计跳过部分 step。梯度累计在半监督任务里很好用做法是先更新多次无监督 loss再统一回传不过这会略微改变 BN 统计分布需要注意。7. 常见问题与调试技巧实录7.1 典型问题速查表我在复现过程中整理了一张速查表按症状、可能原因、解决方法三列来组织遇到问题可以先照着排查症状可能原因解决方法训练初始无监督 loss 恒为 0threshold 太高或模型置信度普遍低下调 threshold 到 0.8 观察覆盖率确认伪标签能产生信号loss_s 正常但准确率不涨RandAugment 强度过大模型在强增强分支学不到有效特征降低 magnitude或删掉 RandomErasing 再试伪标签 mask 覆盖率始终很低训练不充分、模型太弱调高学习率、增加 warmup或先随机初始化训练两个 epoch 再看loss_u为 nansoftmax 后有 0 概率导致 log 爆炸检查是否有 NaN 输入给交叉熵加label_smoothing0.1训练速度极慢no_grad没包住弱增强推理检查弱增强前向是否在torch.no_grad()块内单卡显存溢出unlabeled batch 过大降低mu到 2 或 1或者使用梯度累计复现结果比论文低很多数据增强顺序不正确、没有用 Cutout、随机种子没固定对照官方仓库配置逐项核验先固定随机种子这张表里的问题我都实际遇到过其中“无监督 loss 恒为 0”最容易被忽略因为训练看起来一切正常损失也不炸就是指标上不去。调试时最好在每轮结束时打印mask_ratio能直观看到有效伪标签的占比。7.2 提升训练稳定性的几个关键细节固定随机种子是半监督复现的第一道防线。CIFAR-10 的数据划分、模型初始化、数据 shuffle 三个环节都要固定随机种子否则每次实验的伪标签变化很大调参时根本没法判断改动是不是真的有效。我一般会把torch.manual_seed、np.random.seed和 DataLoader 里的generator都设置好不要省这一步。BatchNorm 的 behavior 在训练和推理阶段不同FixMatch 里无标签 batch 很大BN 统计量会被无标签数据主导。如果 labeled 和 unlabeled 数据分布有偏移这可能会导致问题。经验做法是让 unlabeled batch 和 labeled batch 尽量来自同一个数据域的随机采样这也是为什么论文里都直接对同一个原始数据集划分而不是让 labeled 和 unlabeled 来源不同。最后训练日志里不要只记录 loss 和 acc。至少加上无监督分支的mask_ratio、伪标签置信度均值、学习率三个指标。mask_ratio能反映模型是否在“假装学习”置信度均值能预判模型有多确信自己的预测学习率则能帮助确认 warmup 和 cosine 调度是否按预期执行。有了这三个指标训练过程基本就透明了。我个人在实际操作中的体会是FixMatch 代码看起来非常简洁真正的复杂度全藏在对细节的把控里。同样的算法数据增强顺序写错一位或者无监督 loss 的除法方式选错最终精度能差好几个点。如果你正在做自己的半监督项目先把这一套流程完整跑通再针对数据集特点去调 threshold 和 mu会比盲目堆模型结构高效很多。
返回列表