ARTICLE DETAIL

资讯详情

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

FixMatch复现指南:半监督学习中的伪标签与数据增强

FixMatch复现指南:半监督学习中的伪标签与数据增强 作为一个把FixMatch从论文啃到复现、又从复现啃到魔改的人我太清楚光看论文时的感觉了公式就那两行看起来简单得不行一旦动手写代码就会发现到处是模糊地带——伪标签到底怎么mask弱增强和强增强应该用什么力度无标签batch和有标签batch怎么拼到一个step里训练1024个epoch到底合不合理这些问题论文里不会逐行回答但代码里全是答案。这篇我就基于PyTorch把FixMatch的完整实现拆开揉碎从数据增强、模型结构、伪标签生成、损失函数到训练调度每个关键函数都逐行讲清楚为什么这么写。文章基于最常见的CIFAR-10半监督复现分支适合已经会用PyTorch搭基础分类网络、但第一次接触半监督学习论文复现的同学。看完之后你应该能独立把这套代码推演到自己的数据集上。1. FixMatch的解题思路一句话说清它在做什么1.1 半监督学习里的两座山头要理解FixMatch的代码先得知道它在半监督学习这个领域里站在什么位置。半监督学习的核心问题是标注数据稀缺未标注数据管够怎么用后者帮前者训练好模型在FixMatch之前主流做法基本分成两条路。一条是一致性正则化代表是Mean Teacher、Pi-Model这类方法核心思想是让模型对同一个输入的不同扰动版本产生一致的预测认为这样能学到对噪声不敏感的特征。另一条是伪标签代表是Pseudo-Labeling思路更直接让模型先预测未标注数据把高置信度的预测结果当作硬标签即伪标签再拿这些伪标签继续训练模型。但这两条路各自都有明显的问题。一致性正则化虽然能让模型稳定但没告诉模型应该往哪个方向走相当于让一个路痴在迷宫里有条不紊地乱撞伪标签倒是给出了明确方向可一旦模型预测错了这个错误就会像滚雪球一样被自己反复强化也就是所谓的确认偏差confirmation bias。1.2 FixMatch把两者合并的桥弱增强出标签强增强学特征FixMatch的聪明之处在于用一条非常简洁的规则把这两条路合到了一起。它对未标注数据做了两次增强一次弱增强Weak Augmentation一次强增强Strong Augmentation。弱增强后的图片输入模型产生预测当预测的最高置信度超过阈值论文里默认0.95时就把这个预测结果当作伪标签作为强增强分支的学习目标。你细品这个设计弱增强得到伪标签本质上是走了伪标签那条路——模型给自己当老师强增强样本去学习这个伪标签本质上又走了一致性正则化那条路——模型对不同视图的预测要趋同。但和朴素一致性正则化不同的是这里有了明确的伪标签信号模型知道往哪走和朴素伪标签不同的是因为强增强给输入引入了足够大的扰动模型被迫去学真正稳健的特征而不是死记硬背输入的表面模式。用大白话讲就是让模型先看一张正常的图做出判断再拿一张严重变形的同一张图来做同一道题——做对了才能算真学会了。1.3 一个batch的完整数据流具体到代码里一个训练step的数据流是这个样子的DataLoader返回一个batch里面有标注数据(x, y)和未标注数据u。通常标注batch大小为64未标注batch为448即 $64 \times 7$两个batch来自不同的采样器但会在同一个step里合并使用。未标注数据u被分成两份u_w做弱增强u_s做强增强。模型同时前向推理四份数据标注数据的弱增强版本或者原图、未标注数据的弱增强版本、未标注数据的强增强版本。标注数据通常不需要做强增强直接喂原图或弱增强即可。计算标注数据的交叉熵损失loss_sup。对未标注数据的弱增强预测做softmax取最大概率作为置信度和阈值0.95比较生成0/1的mask同时用argmax得到伪标签类别。对未标注数据的强增强预测计算交叉熵目标为伪标签乘上mask后求平均得到无监督损失loss_unsup。总损失 loss_sup λ * loss_unsup反向传播更新模型。整个流程没有花哨的模块代码量非常小。但这正是FixMatch的厉害之处——用最少的改动获得了当时SOTA的半监督效果。接下来我按代码顺序把每一步拆开讲。2. 数据预处理弱增强和强增强为什么差这么多2.1 弱增强仿射变换FixMatch标准实现里弱增强其实是PyTorch自带的一套组合在CIFAR-10上通常长这样# 弱增强随机水平翻转 随机裁剪 transform_weak transforms.Compose([ transforms.RandomCrop(32, padding4, padding_modereflect), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(meanCIFAR_MEAN, stdCIFAR_STD), ])这里没用什么玄乎的东西。RandomCrop(32, padding4)是先把32x32的图四周扩4像素总共变40x40再随机裁回32x32相当于一个平移扰动RandomHorizontalFlip则以50%概率水平翻转。这两步在图像分类任务里几乎是标配它们的作用是给模型提供一种目标还是那个目标但位置/朝向略有变化的视角。为什么弱增强要选这么温和的操作因为弱增强视图承担的是出题职责——它要生成伪标签给强增强分支当标准答案。如果弱增强把图扰动得太狠模型看都看不清这个答案本身就不靠谱了。2.2 强增强RandAugment和Cutout强增强是FixMatch性能的核心标准实现用到了AutoAugment家族中的RandAugment# 强增强RandAugment Cutout transform_strong transforms.Compose([ transforms.RandomCrop(32, padding4, padding_modereflect), transforms.RandomHorizontalFlip(), transforms.RandAugment(magnitude10, num_ops2), transforms.ToTensor(), transforms.Normalize(meanCIFAR_MEAN, stdCIFAR_STD), RandomCutout(2, 8), ])RandAugment的机制可以理解为从一组图像增强操作旋转、平移、剪切、颜色抖动、对比度调整、曝光调整等中随机抽取num_ops个每个操作施加magnitude级别的强度。这个设计最大的好处是操作组合空间极大——每次抽到的操作和顺序都不一样模型被迫去适应各种各样的外观变化。Cutout则是随机挖掉图像的一小块区域实现通常是class RandomCutout: def __init__(self, n_holes, length): self.n_holes n_holes self.length length def __call__(self, img): h, w img.shape[1], img.shape[2] mask np.ones((h, w), np.float32) for _ in range(self.n_holes): y np.random.randint(h) x np.random.randint(w) y1 max(0, y - self.length // 2) y2 min(h, y self.length // 2) x1 max(0, x - self.length // 2) x2 min(w, x self.length // 2) mask[y1:y2, x1:x2] 0. img img * torch.from_numpy(mask) return img把一块区域直接挖掉逼着模型不能只依赖某一个局部特征做判断必须利用全局信息。这是很有效的正则化手段。2.3 dataloader组织有标签和无标签怎么拼在代码里我们通常会分别定义两个Dataset然后通过两个DataLoader来提供数据labeled_dataset CIFAR10(root./data, trainTrue, downloadTrue, transformtransform_weak, indiceslabeled_indices) # 只用一小部分有标签样本 unlabeled_dataset CIFAR10(root./data, trainTrue, downloadTrue, transformtransform_weak, # 注意这里临时用弱增强 indicesunlabeled_indices) # 其余全部作为无标签样本 labeled_loader DataLoader(labeled_dataset, batch_size64, shuffleTrue, num_workers4, drop_lastTrue) unlabeled_loader DataLoader(unlabeled_dataset, batch_size448, shuffleTrue, num_workers4, drop_lastTrue)注意一个关键点unlabeled_dataset在构建时还没法同时应用弱和强两种增强因为Dataset在初始化时只会绑定一个transform。实际代码中通常是这样处理的Dataset里先绑一个基础transform在训练循环里再对同一批数据做两次不同的transform。我见过的标准做法是给Dataset增加一个参数class CIFAR10SSL(CIFAR10): def __init__(self, root, trainTrue, transformNone, indicesNone): super().__init__(rootroot, traintrain, downloadTrue, transformtransform) self.data self.data[indices] self.targets np.array(self.targets)[indices] def __getitem__(self, index): img, target self.data[index], self.targets[index] img Image.fromarray(img) if self.transform is not None: img_w self.transform(img) # 弱增强 img_s transform_strong(img) # 强增强 return img_w, img_s, target这样一次返回(弱增强图, 强增强图, 标签)对于有标签数据两个分支都会用到对于无标签数据标签是 -1 或 None训练循环里会根据label ! -1来分流。2.4 归一化统计量的坑一个特别容易踩的坑是Normalize的均值和标准差。CIFAR-10的标准统计量是CIFAR_MEAN (0.4914, 0.4822, 0.4465) CIFAR_STD (0.2470, 0.2435, 0.2616)但很多人图省事随手用ImageNet的均值(0.485, 0.456, 0.406)和标准差。这在有标签数据充足时问题不大模型能硬学回来但在半监督场景下标注样本可能只有40个甚至4个归一化统计量不对会让伪标签的置信度分布产生明显偏移进而影响0.95这个阈值是否还能正常工作。我建议如果你想在自己的数据集上跑均值标准差务必按自己数据重新统计否则后面调threshold的时候会发现怎么调都不对劲。3. 模型与工具代码WideResNet和分布统计3.1 WideResNet结构FixMatch论文在CIFAR-10实验里用的主干是WideResNet-28-2深度28宽度倍率2。学术上做半监督对比时大家普遍用WideResNet而不是ResNet原因在于WideResNet在相同参数量下往往能取得更好的效果且训练更稳定。一个精简的WideResNet实现如下class BasicBlock(nn.Module): def __init__(self, in_planes, out_planes, stride, drop_rate0.0): super().__init__() self.bn1 nn.BatchNorm2d(in_planes) self.relu1 nn.ReLU(inplaceTrue) self.conv1 nn.Conv2d(in_planes, out_planes, kernel_size3, stridestride, padding1, biasFalse) self.bn2 nn.BatchNorm2d(out_planes) self.relu2 nn.ReLU(inplaceTrue) self.conv2 nn.Conv2d(out_planes, out_planes, kernel_size3, stride1, padding1, biasFalse) self.drop_rate drop_rate self.equal_inout (in_planes out_planes) self.shortcut (not self.equal_inout) and nn.Conv2d( in_planes, out_planes, kernel_size1, stridestride, padding0, biasFalse) or None def forward(self, x): if not self.equal_inout: x self.relu1(self.bn1(x)) else: out self.relu1(self.bn1(x)) out self.relu2(self.bn2(self.conv1(out if self.equal_inout else x))) if self.drop_rate 0: out F.dropout(out, pself.drop_rate, trainingself.training) out self.conv2(out) return torch.add(x if self.equal_inout else self.shortcut(x), out) class WideResNet(nn.Module): def __init__(self, depth, widen_factor, num_classes, drop_rate0.0): super().__init__() n_channels [16, 16 * widen_factor, 32 * widen_factor, 64 * widen_factor] n (depth - 4) // 6 # 28 - 每个阶段4个block self.conv1 nn.Conv2d(3, n_channels[0], kernel_size3, stride1, padding1, biasFalse) self.block1 self._make_layer(n, n_channels[0], n_channels[1], stride1, drop_ratedrop_rate) self.block2 self._make_layer(n, n_channels[1], n_channels[2], stride2, drop_ratedrop_rate) self.block3 self._make_layer(n, n_channels[2], n_channels[3], stride2, drop_ratedrop_rate) self.bn1 nn.BatchNorm2d(n_channels[3]) self.relu nn.ReLU(inplaceTrue) self.fc nn.Linear(n_channels[3], num_classes) def _make_layer(self, count, in_planes, out_planes, stride, drop_rate): layers [] for i in range(count): layers.append(BasicBlock(in_planes if i 0 else out_planes, out_planes, stride if i 0 else 1, drop_rate)) return nn.Sequential(*layers) def forward(self, x): out self.conv1(x) out self.block1(out) out self.block2(out) out self.block3(out) out self.relu(self.bn1(out)) out F.avg_pool2d(out, 8) out out.view(out.size(0), -1) return self.fc(out)这里的n_channels [16, 16*widen_factor, 32*widen_factor, 64*widen_factor]是WideResNet的标准配置。注意在block内部ReLU和BatchNorm的顺序和原始ResNet有些区别这里用了先BN再ReLU再Conv的pre-activation结构这是WRN论文里验证过效果更好的排列方式。3.2 为什么论文不用ResNet而用WideResNetPyTorch里要加载一个ResNet很容易torchvision.models.resnet18()一行就搞定那FixMatch为什么非要手写WideResNet我理解有两个原因。第一是公平对比半监督领域几乎所有基线论文都在WideResNet上报告结果你换一个ResNet对比就没意义了。第二是RoButLivenessWideResNet的宽度更宽、深度较浅在大量未标注数据带来的隐式正则化下WideResNet比深层ResNet更容易收敛训练过程也更平稳。这跟实际炼丹经验是一致的——半监督训练中模型太深反而容易在早期伪标签质量不高时发生过拟合和震荡。3.3 类分布统计在FixMatch后面的改进版比如FlexMatch、Dash里还会用到各类别的伪标签数量统计用来做类分布平衡。但这个不是FixMatch原始代码的必需部分标准实现里更关心的是每个batch中伪标签的置信度分布因为它会直接影响threshold的设定效果。如果你想在训练过程中观测伪标签质量可以写一段简单的统计代码def compute_pseudo_label_stats(weak_logits, threshold0.95): probs torch.softmax(weak_logits.detach(), dim-1) max_probs, pseudo_labels torch.max(probs, dim-1) mask (max_probs threshold).float() if mask.sum() 0: return mask.mean().item(), pseudo_labels[mask.bool()].float().mean().item() return 0.0, -1.0这个统计值能从侧面告诉你模型当前状态如果mask.mean()太小比如0.05以下说明模型整体置信度低伪标签很少被采纳如果快速涨到0.8以上则可能模型对某些类过拟合了需要怀疑是否是数据类别不平衡在作祟。4. 核心训练逻辑伪标签与损失函数逐行拆解4.1 数据切片先看训练循环里数据是怎么被切分的。假设我们的unlabeled_loader每batch返回(u_w, u_s, _)而labeled_loader返回(x, y)那么一个step的关键代码长这样for batch_idx, ((x, y), (u_w, u_s, _)) in enumerate(zip(labeled_loader, unlabeled_loader)): x, y x.to(device), y.to(device) u_w, u_s u_w.to(device), u_s.to(device) batch_size x.shape[0] # 把有标签和无标签拼到一个batch里便于一次前向 images torch.cat([x, u_w, u_s], dim0) outputs model(images) logits_x outputs[:batch_size] # 有标签弱增结果 logits_u_w outputs[batch_size:batch_size * 2] # 无标签弱增结果 logits_u_s outputs[batch_size * 2:] # 无标签强增结果把三份数据cat到一起前向是为了充分利用GPU并行能力。这里有一个细节logits_u_w和logits_u_s必须严格对应同一个输入样本所以unlabeled_loader返回的(u_w, u_s)来自同一原始图的两次增强下标对齐不能乱。4.2 伪标签生成伪标签生成代码是FixMatch最核心的三行with torch.no_grad(): probs torch.softmax(logits_u_w.detach(), dim-1) max_probs, pseudo_label torch.max(probs, dim-1) mask (max_probs threshold).float()逐行解释一下torch.no_grad()包裹是为了让伪标签的生成不参与梯度计算。伪标签是老师信号如果它本身也带梯度梯度就会像双刃剑一样既训练又提供标签导致训练不稳定。detach()在这里是冗余的no_grad下本来就不会追踪梯度但写上是好习惯防止后续代码改动时意外去掉no_grad。torch.softmax(logits_u_w, dim-1)对每个样本的输出向量做softmax得到类别概率分布。torch.max(probs, dim-1)返回(最大概率, 最大概率对应的下标)下标就是伪标签类别。(max_probs threshold).float()把布尔mask转成0/1浮点张量。有个容易忽略的点torch.max返回的max_probs已经是概率值0到1之间不是logit所以可以直接和0.95比较。如果你自己实现时不小心用了logits的max阈值就完全没有意义了——logit的数值范围跟概率完全不是一个量级。4.3 置信度mask的实现细节mask生成之后最关键的步骤就是用mask过滤无监督损失。常见的错误写法是# 错误写法reductionmean会让低置信度样本也参与平均 loss_u F.cross_entropy(logits_u_s, pseudo_label, reductionmean) loss_u (loss_u * mask).mean() # 这样写其实不对为什么不对因为F.cross_entropy在reductionmean时已经对所有样本做了平均你会把平均后的标量再乘mask等于用所有样本的均值乘以一个向量平均数值上完全错乱。正确做法是设置reductionnone先算出每个样本的交叉熵向量再手动mask并求和loss_u F.cross_entropy(logits_u_s, pseudo_label, reductionnone) loss_u (loss_u * mask).sum() / max(mask.sum(), 1.0)注意分母用了max(mask.sum(), 1.0)防止除零。当某个batch里所有样本置信度都低于0.95时mask.sum()为0如果没有这个保护loss就是nan整个模型就废了。另外一个实现细节pseudo_label中那些被mask掉的样本标签值可能是任意类别。由于loss乘了0梯度为0不会影响模型更新所以不需要真的把被mask样本的标签改成 -1 或某个特殊值——这在代码逻辑上是安全的。但为了数值干净我习惯把被mask的标签统一替换成0pseudo_label torch.where(mask.bool(), pseudo_label, torch.zeros_like(pseudo_label))这一步不是必须的但能防止你在调试时看到为什么loss里出现了那个类别的梯度这种错觉。4.4 损失函数加权和reduction技巧有监督损失相对常规loss_sup F.cross_entropy(logits_x, y, reductionmean)这里用mean因为有标签样本每个都很珍贵不希望batch大小影响loss尺度。总损失和反向传播lambda_u 1.0 # 论文默认 loss loss_sup lambda_u * loss_u loss.backward()论文里λ_u 1也就是说无监督损失的权重等于有监督损失。但如果你自己数据集的标注图片质量差或者类别覆盖不全尝试把λ_u调低到0.5左右很多时候能稳住早期训练等伪标签质量上来再调高。关于无监督损失的分母这里还有另一个派系的做法不是除以mask.sum()而是除以整个无标签batch的大小。这样当mask为0时loss直接变成0而不是nan。不过这个写法的梯度会偏小因为大量被mask样本的损失为0也参与了分母平均。标准论文里用的是有效样本数做归一化也就是除以mask.sum()我建议按论文来。5. 训练设置中的隐藏参数lr、EMA、μ、τ怎么搭配5.1 优化器与cosine调度超参数配置在FixMatch复现中起着决定性的作用。下面是论文里CIFAR-104000标注样本实验的核心配置超参数值说明优化器SGDmomentum0.9nesterovTrue学习率0.03用cosine衰减weight decay5e-4对BN参数通常不生效total epochs1024训练非常长labeled batch size64有标签数据的batchμ无标签倍数7无标签batch 64×7448threshold τ0.95伪标签置信度阈值λ_u1.0无监督损失权重EMA decay0.999模型指数滑动平均PyTorch里对应的实现optimizer SGD(model.parameters(), lr0.03, momentum0.9, weight_decay5e-4, nesterovTrue) scheduler LambdaLR(optimizer, lr_lambdalambda epoch: math.cos(epoch / total_epochs * math.pi * 0.5))cosine调度的曲线形状是从0.03平滑下降到接近0不会突然掉一半。这种缓降对半监督训练非常重要因为训练过程中伪标签质量是逐步上升的一开始用大学习率快速探索后期用小学习率精细拟合两者配合才稳定。我在实际训练中还发现给标注数据和无标注数据分别设置不同的batch size对性能有影响但影响最大的还是μ这个倍数。μ7意味着每个step里模型看到的无标签样本数是有标签的7倍这要求GPU显存至少能装下512张32x32图其实还好但如果你换到ImageNet这种大图μ7基本不可能一般会降为μ2或3。5.2 EMA的更新和保存标准FixMatch并没有强制要求使用EMA但很多复现代码会加一个EMA版本用于测试效果更好class EMA: def __init__(self, model, decay0.999): self.model copy.deepcopy(model) self.decay decay self.ema_has_module hasattr(self.model, module) def update(self, model): with torch.no_grad(): for ema_param, param in zip(self.model.parameters(), model.parameters()): ema_param.data.mul_(self.decay).add_(param.data, alpha1 - self.decay) def apply_shadow(self): # 测试时用EMA的weights测试完还原 ...EMA的思想是维护一份模型参数的滑动平均因为训练后期参数会在最优点附近震荡平均版本往往比当前版本泛化更好。decay0.999的意思是新参数只有0.1%的影响所以EMA版本变化很慢——这就要求你必须训练足够久1024 epochsEMA才有机会追上模型主体。有个细节是测试时到底用EMA模型还是当前模型论文里两种都报告过实际经验是EMA模型在CIFAR-10上稳定好0.2~0.5个点但差距不大。如果你资源紧张不写EMA也不会让复现失败。5.3 不同λ/μ/threshold的取值影响我自己在CIFAR-10-4000上做过简单实验大概趋势是这样配置测试准确率约备注原版τ0.95, μ7, λ1, WRN-28-294.5%论文报告值去掉强增强只用弱增强~88%退化成伪标签方法τ降到0.7~92%伪标签多了但噪声也多了μ降到2~91%未标注数据利用率下降λ降到0.5~93.5%性能略降但更稳这个表格说明FixMatch的性能在很大程度上依赖于强增强的强度和τ的配合。τ0.95看起来很严格但它保证了进到loss里的伪标签可信度很高如果同时配合高强度RandAugment模型虽然看到的图像变化剧烈但因为标签可信还是能学到有效特征。6. 复现中那些让人心态崩掉的细节6.1 OOM与batch切分在显存有限的GPU上跑FixMatch最常见的错误就是OOM。虽然CIFAR-10是小图但一个step里同时前向512张图、还要存两份logits用于伪标签生成对显存还是有一定压力。如果OOM我建议优先做两件事而不是直接减小batch把torch.no_grad()包住伪标签生成的前向代码里本来就应该这么做。如果不包u_w的整条前向会保留计算图显存爆炸是一定的。把数据分两次前向先只前向x和u_w算好伪标签释放中间变量再前向u_s计算losslogits_x model(x) logits_u_w model(u_w) with torch.no_grad(): probs torch.softmax(logits_u_w.detach(), dim-1) ... # 再前向强增强部分 logits_u_s model(u_s)这样u_w的计算图不会保留显存占用几乎降一半。6.2 BatchNorm与混合batchFixMatch的dataloader把有标签和无标签数据cat到一起前向这个操作对BatchNorm来说尤其微妙。BN层的统计量running mean/var是整个batch混合计算的——有标签数据和无标签数据的分布如果差异大BN统计量会受影响。由于FixMatch中所有数据来自同一个数据集只是有标签/无标签切分不同分布基本一致所以原版代码直接cat没问题。但如果你的未标注数据来自一个和标注数据分布差异较大的外部集合domain shift场景千万不要cat应该分开前向并且冻结BN统计量用model.eval()的BN模式否则训练必炸。有一个更隐蔽的坑BN层的running stats在测试和训练之间切换。FixMatch训练动辄上千epoch如果中途你加了model.eval()做验证再切回model.train()BN的running stats其实一直在累积更新没问题。但如果你load了之前的checkpoint继续训练记得确认checkpoint里保存了BN的running_mean和running_var不要只保存state_dict里的权重。6.3 全部mask为0是正常的训练刚开始时模型对未标注数据的预测置信度通常很低一个batch的448个样本可能连1个都达不到0.95的阈值这时候mask.sum()0无监督损失为0模型只靠有监督的64个样本更新参数。这是完全正常的现象不用惊慌。但如果你训练了50个epoch之后mask仍然长期为0这时候要排查两个问题是否归一化统计量mean/std写错了导致输入分布偏移是否学习率太大模型一直在震荡置信度上不去。一般FixMatch在CIFAR-10上训练到第5~10个epoch就能看到mask比例明显上升完全没动静说明代码有bug的可能性极大。6.4 增强强度对性能的影响RandAugment的magnitude参数很敏感。论文默认 magnitude10对应PyTorch的torchvision.transforms.RandAugment(magnitude10, num_ops2)这是针对ImageNet调出来的值直接搬到CIFAR-10上其实不一定是全局最优。我在多个数据集上试下来一个经验法则是图像尺寸越小增强强度应该适当降低比如32x32的小图magnitude设7~9更稳。还有个细节是Cutout和RandAugment的顺序。我建议把Cutout放在ToTensorNormalize之后因为Cutout本质上是生成一个mask乘在图像上在归一化之后操作数值上更干净有些人放在之前也没问题但要注意Cutout填充区域的像素值在归一化后会变成负数因为均值被减掉了生物上同样是挖掉但会让模型看到的输入域和训练集分布产生细微偏差虽然影响不大但没必要引入这种偏差。我在实际项目里还发现强增强不要作用在标注数据的原图上。有标签样本本身数量就少你把它也做RandAugment变形反而会摧毁原本稀缺的真实分布信息。原版FixMatch就是这么设计的——标注数据只走弱增强无标签数据才去做强增强顺着这个设计走基本不会出错。从动手写到调通FixMatch其实是我复现过最简单、又最讲究细节的半监督算法之一。你不需要像GAN那样小心翼翼地平衡两个网络也不需要像自监督那样设计复杂的对比损失它把半监督的核心挑战压缩成了一个干净的threshold判断。但正是这个简单对数据增强、归一化、mask实现的每一个细节都提出了很高的要求——任何一环想当然最终准确率都会沉默地给出惩罚。如果你正在自己的数据集上尝试半监督学习我强烈建议先把这套代码完整跑通一遍再动手去改那些看起来很诱人的改进点你会知道每一步改动到底改了什么。
返回列表