拓冰建站拓冰建站
首页 / 资讯中心 / 正文

MoBY核心机制与源码解析:ViT自监督学习的关键工程实践

简介面向自监督学习与深度学习研究者的MoBY完整源码包以Vision Transformer为主干架构将MoCo v2和BYOL的设计思想有机结合在ImageNet-1K线性评估中经过300轮训练后DeiT-S和Swin-T分别取得72.8%与75.0%的top-1准确率并且相比MoCo v3、DINO所需的训练技巧更加轻量实用。压缩包共33个文件以Python源码、YAML配置、Markdown文档和PNG示意图为主整体大小仅1.11MB其中Python脚本覆盖数据加载、模型构建、训练与线性评估等模块YAML配置则提供不同backbone和训练策略的组合。目前已有516人学习下载适合希望快速复现MoBY结果、理解自监督对比学习流程或在此基础上进行改进的研究者与工程师。包内附模型定义、训练主程序、配置体系、快速开始文档和架构示意图能够完整支撑从环境搭建、实验复现到后续改进的全过程帮助读者深入掌握MoBY的核心实现与轻量trick。1. 为什么自监督学习转向 Vision Transformer 时MoBY 是绕不开的一个点自监督学习在 2021 年前后迎来了一个转折主流方法集体从卷积网络转向 Vision Transformer。MoBY 恰好站在这个交叉点上它把 MoCo v2 的对比学习框架和 BYOL 的非对称结构合并到同一套流程里让 ViT 在没有 ImageNet 标签的情况下获得可迁移的表征。标题里的「数据源码」暗示了它的定位——不是一篇只看结论的论文速览而是值得拿到本地对照源码复现的工程样板。你最好已经跑过一遍基础的对比学习训练比如 SimCLR 或 MoCo v2没有的话也建议先把深度学习环境配置跑通再往下读。下面从它的核心机制开始逐步讲到数据加载、动量更新和损失计算的写法以及我实际调参时遇到的坑。2. MoBY 的核心机制动量编码器、预测器与不对称设计先澄清一个容易搜错的地方自监督学习里的 MoBY 和 Docker 生态里那个 Moby 没有任何关系检索时如果混入「dockerd」「容器」关键词会带偏方向。这里的 MoBY 是 2021 年底公开的自监督学习方法全称是 Self-Supervised Learning at Scale来自工业界实验室论文方向是在大规模数据集上做无标签预训练。2.1 从 MoCo v2 与 BYOL 看 MoBY 的定位MoBY 不是凭空造出来的新范式而是把 2021 年前后两个有代表性的方法做了结构性合并。MoCo v2 提供了对比损失和动量队列query encoder 看到一张图的增强视图key encoder 看到同一个图的另一个增强视图负样本来自队列损失用 InfoNCE。BYOL 则提供了 predictor 和非对称更新方式online 分支末尾接一个 predictor训练目标是让 predictor 的输出逼近 target 分支的表示整个过程不依赖负样本。MoBY 的做法是保留 MoCo 的队列与 InfoNCE 损失同时加入 BYOL 的 predictor并让 query encoder 这条支路同时承担两个任务——对比损失和自蒸馏损失。这样设计的原因要从 ViT 的习性说起。ResNet 靠卷积的局部归纳偏置哪怕增强比较重也能稳定收敛ViT 没有这种先验直接套用 MoCo v2 的训练配置会发现模型很难收敛或者精度明显低于同代的卷积方法。MoBY 把 predictor 放在 query 支路本质上是在「学生」上加了一个可学习的小网络迫使表示空间更平滑同时用 key 编码器作为缓慢移动的「老师」缓解 ViT 对小扰动敏感的毛病。2.2 动量编码器与 predictor 在 MoBY 里的不对称设计结构上 MoBY 有两条编码器结构完全相同都是 ViT。一条是 query encoder学生一条是 key encoder老师。query encoder 后面额外接一个 predictorkey encoder 后面是干净的表示。前向时同一张图的两个全局视图 x1、x2 各自独立地通过 query encoder 得到 q1、q2再通过 predictor 得到 z1、z2这两个视图也会通过 key encoder 得到 k1、k2。其余的小裁剪局部视图只走 key encoder不参与对比损失只参与后续的蒸馏损失。动量更新的公式是 θ_k ← m · θ_k (1 − m) · θ_q。参数 m 不是固定的而是从 0.99 线性增加到 1.0。这里有个容易被忽略的细节在 MoCo v2 里 m 通常取常数 0.999而 MoBY 选择了 schedule 形式的动量。原因是 ViT 的训练对「老师」的更新幅度更敏感前期 m 较小让老师尽快跟上学生的变化后期 m 无限接近 1.0 相当于把老师冻结输出趋于稳定蒸馏损失才能提供可靠的 soft target。我一般会关注 m 是否在 scheduler 里被正确更新因为很多复现只改了数值而忘了写 schedule。2.2.1 动量更新在代码里的写法动量更新的代码很短但顺序和写法决定训练是否稳定def update_key_encoder(query_net, key_net, m): for param_q, param_k in zip(query_net.parameters(), key_net.parameters()): param_k.data m * param_k.data (1.0 - m) * param_q.data这段代码必须在optimizer.step()之后调用。如果放在 backward 之前key encoder 会用到未更新的学生权重动量更新就失去了「用新学到的知识缓慢修正老师」的意义。param_k.data直接做原地赋值不通过copy_或assign是为了避免破坏 PyTorch 的自动求导图。key encoder 的参数在初始化时设置了requires_grad False所以这里不担心梯度回传。动量参数怎么配直接决定 ViT 训练曲线长什么样。下面是我在做消融时常用的配置对照配置动量策略对 ViT 的典型影响MoCo v2 风格固定 0.999教师更新慢队列稳定但 ViT 容易欠拟合BYOL 风格固定 0.99教师跟随快ResNet 上稳定ViT 上精度波动大MoBY 默认0.99 线性升到 1.0教师前期快速适应、后期冻结ViT 训练最稳2.3 多尺度裁剪与两种损失的配合MoBY 在输入侧采用了 BYOL 中的 multi-crop 策略一张图先生成 2 个全局裁剪尺寸是原图的 224×224再生成 8 个局部裁剪尺寸是 96×96。全局裁剪负责语义局部裁剪负责细节和尺度不变性。训练时只有全局裁剪之间的两两对比进入 InfoNCE 损失蒸馏损失则把每个局部视图的表征压向同一张图的全局视图表征。这样设计的收益是不增加对比损失的计算量就能利用更丰富的像素信息。损失函数两个分量的组合也很关键。对比损失沿用 InfoNCE公式是 L_con −log exp(q·k/τ) / Σ exp(q·k/τ)其中 q 来自 predictor 输出k 来自 key 编码器负样本从动量队列里取。蒸馏损失则是把 key 编码器对局部视图的输出当作目标让 query 编码器对全局视图的输出去回归它。为什么需要这个分量因为对比损失只约束全局视图之间的判别性对局部和全局的尺度不变性没有显式约束而关键的语义往往需要跨越尺度才能学到。3. MoBY 源码的工程骨架从配置到训练循环3.1 配置文件里的关键参数无论是读源码还是自己复现我一般会从配置文件入手先建立对训练流程的全局印象。MoBY 这类自监督学习的配置项比普通分类模型多因为除了模型结构还要管队列、动量、多裁剪增强和双损失。下表是我读 MoBY 源码时梳理出的关键参数参数常见取值作用encoderViT-B/16backbone可换成 ViT-L/16local_crops_number8局部裁剪数量global_crops_size224全局裁剪边长local_crops_size96局部裁剪边长momentum_start0.99EMA 动量初值momentum_end1.0EMA 动量终值batch_size2048~4096全局 batch多卡累计lr0.0003 × (batch_size/2048)学习率线性缩放optimizerLARS对大 batch 更稳的优化器loss_weights对比 1.0 蒸馏 1.0两个损失的比例读配置的时候我会特别关注 batch_size 和队列长度的匹配。InfoNCE 的负样本由队列提供队列长度如果远大于 batch_size表示对比目标更硬训练更慢但表征更稳健如果队列只比 batch_size 大一点等价于只用 in-batch 负样本InfoNCE 的优势就消失了。3.2 模型定义encoder 与 predictor 的组装下面这段是简化后的 PyTorch 风格代码展示 MoBY 的模型骨架。核心是两个共享结构但参数独立的 encoder以及挂在 query 分支上的 predictorimport torch import torch.nn as nn class MoBY(nn.Module): def __init__(self, dim768, momentum0.99): super().__init__() import timm # 实际工程中加载 ViT-B/16并把分类头替换为 dim 维投影输出 self.query_encoder timm.create_model(vit_base_patch16_224, num_classesdim) self.key_encoder timm.create_model(vit_base_patch16_224, num_classesdim) # predictor 只挂在 query 分支key 分支没有 self.predictor nn.Sequential( nn.Linear(dim, dim * 4, biasFalse), nn.BatchNorm1d(dim * 4), nn.ReLU(inplaceTrue), nn.Linear(dim * 4, dim) ) # 初始化 key 编码器与 query 编码器相同并冻结参数 for param_q, param_k in zip(self.query_encoder.parameters(), self.key_encoder.parameters()): param_k.data.copy_(param_q.data) param_k.requires_grad False self.momentum momentum这里有几个参数层面的说明。momentum是 EMA 的基线值实际训练中会按 schedule 传入新的 m而不是始终用初始化值。requires_grad False保证了 key 编码器只通过动量更新改变权重不接收梯度。predictor 中间层用 BatchNorm1d 而不是 LayerNorm是沿用 BYOL 的设置如果换成别的归一化发现精度掉点先不要怀疑优化器先改回 BatchNorm1d 再对比。3.3 训练循环里的动量更新与损失计算训练循环的关键顺序是先算 loss再更新梯度最后做动量更新。顺序不能反因为动量更新要使用刚刚更新完的学生参数for x1, x2, x_local in loader: # x1/x2 是全局裁剪x_local 是局部裁剪 z1 model.predictor(model.query_encoder(x1)) z2 model.predictor(model.query_encoder(x2)) with torch.no_grad(): k1 model.key_encoder(x1) k2 model.key_encoder(x2) k_local model.key_encoder(x_local) loss_con contrastive_loss(z1, k2, queue) contrastive_loss(z2, k1, queue) loss_dis distill_loss(z1, k_local) distill_loss(z2, k_local) loss loss_con loss_dis loss.backward() optimizer.step() m momentum_schedule(epoch, total_epochs) update_key_encoder(model.query_encoder, model.key_encoder, m) queue.enqueue(torch.cat([k1, k2], dim0))代码里with torch.no_grad()包裹 key encoder 的前向避免梯度从 key 分支回流。queue是一个先进先出的张量队列每个 step 后把新的 key 入队把最老的 key 出队。momentum_schedule通常返回一个随 epoch 线性变化的浮点数最后一轮接近 1.0。常见误用是忘记把 key encoder 设为 eval 模式或忘记关闭梯度另一种是在算蒸馏损失时把 k1、k2 和局部视图混在一个 batch 里导致梯度流向 key 分支。遇到这种情况loss 和精度都正常但显存占用明显偏高、训练速度变慢最直接的排查方法就是检查k1.requires_grad。4. 数据加载与 augmentation 在 MoBY 里的具体落法4.1 同一 batch 里两种视图分别喂给哪条编码器MoBY 的数据组织方式决定了它的效率。每个训练样本要生成 10 个裁剪2 个全局 8 个局部。在 DataLoader 里为了省内存常见做法是把全局裁剪和局部裁剪分别拼成独立的 batch。也就是说每个 iteration 返回(x1, x2, x_local)其中 x1、x2 的形状是[B, 3, 224, 224]x_local 的形状是[B * 8, 3, 96, 96]。x1 和 x2 各自通过 query encoder 和 key encoderx_local 只通过 key encoder。这个设计的边界条件在于对比损失需要的是「同一张图的两个全局视图互为正样本」因此在组织 batch 时要保证x1[i]和x2[i]来自同一个样本。如果我直接用一个随机采样器生成两个独立 tensor 而不管它们是否配对那么在计算 InfoNCE 时正样本对就错位了。复现时一定要先做一个 sanity check打印x1[i]和x2[i]的索引确认来自同一个原图。4.2 Augmentation 参数与代码实现MoBY 沿用了 BYOL 的增强策略包括随机裁剪、颜色抖动、灰度化、高斯模糊和太阳能化。全局裁剪的面积比例是 0.14~1.0局部裁剪是 0.05~0.14。这个面积区间的区别很重要全局裁剪要保留物体的主体局部裁剪只保留一部分细节目的是让模型学会大尺度和局部信息的一致性。用torchvision.transforms实现时我会先定义一个基础增强函数再分别组合出全局和局部两套from torchvision import transforms def get_moby_transforms(): def global_view(): return transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.14, 1.0)), transforms.RandomHorizontalFlip(), transforms.RandomApply( [transforms.ColorJitter(0.4, 0.4, 0.4, 0.1)], p0.8), transforms.RandomGrayscale(p0.2), transforms.RandomApply( [transforms.GaussianBlur(23, sigma(0.1, 2.0))], p1.0), transforms.RandomSolarize(threshold0.5, p0.2), ]) def local_view(): return transforms.Compose([ transforms.RandomResizedCrop(96, scale(0.05, 0.14)), transforms.RandomHorizontalFlip(), transforms.RandomApply( [transforms.ColorJitter(0.4, 0.4, 0.4, 0.1)], p0.8), transforms.RandomGrayscale(p0.2), transforms.RandomApply( [transforms.GaussianBlur(23, sigma(0.1, 2.0))], p0.5), ]) return global_view, local_view参数说明RandomResizedCrop的第一个参数是输出尺寸第二个参数scale控制裁剪面积占原图的比例。GaussianBlur的核大小取 23sigma 范围 0.1~2.0局部视图的模糊概率降到 0.5。这些数值直接影响对比任务和蒸馏任务的难度改的时候不要一次性改多个变量否则无法定位是哪个增强导致精度变化。4.3 数据加载时序的坑增强顺序与多卡采样RandomResizedCrop与RandomSolarize的组合顺序会影响结果。太阳能化应当在颜色抖动之后、归一化之前。假如把归一化放在增强中间像素值被缩放到 0~1 附近太阳能化的阈值会失配导致大量样本被错误地压到 0表征质量明显下滑。这类问题不报错只能靠验证集精度变化来发现。当 batch_size 达到 2048 以上时单卡显存放不下必然会用多卡。多卡场景下每个 rank 各自加载一个子 batch队列在每张卡上各维护一份还是全局共享会直接影响负样本多样性。MoBY 类方法通常让每张卡维护自己的队列不做卡间同步因为队列长度本身已经足够大卡间多样性提升有限。复现时如果发现 loss 在不同卡之间差距很大优先怀疑 DataLoader 的 shuffle 逻辑和随机种子设置而不是队列同步。5. 复现 MoBY 时最值得调的 3 个参数5.1 学习率与 batch size 的线性关系MoBY 的学习率并不是直接照搬配置就能跑通的。常见设置是lr 0.0003 × (batch_size / 2048)。如果你只把 batch_size 从 2048 提到 4096而不把学习率翻倍训练 loss 会显得更平滑但最终的线性评估精度会略低反过来学习率追太高ViT 的 attention 权重会在前几个 epoch 就发散。判断当前学习率是否合适的经验是前 20 个 epoch 里 InfoNCE 的 loss 应当持续下降如果出现波浪形震荡先降到 1/3 再试。5.2 动量系数的 schedule 与 ViT 的兼容性动量系数从 0.99 线性升到 1.0 这条曲线是 MoBY 针对 ViT 做的关键调整。如果你把 m 固定成 0.999训练曲线看起来差不多但最终精度会掉 1 到 2 个点如果固定成 0.99key encoder 更新过快蒸馏损失会失去稳定目标。我一般会在验证集上画出 m 随时间变化的曲线与精度曲线对齐确认在训练前 60% 阶段 m 还在上升而不是一上来就接近 1.0。这个 schedule 用几行代码就能实现def momentum_schedule(epoch, total_epochs): # 返回当前轮次的动量系数线性从 0.99 增长到 1.0 return 0.99 (1.0 - 0.99) * epoch / total_epochs5.3 局部裁剪数量与显存的取舍把局部裁剪数量从 8 减到 4显存能下降约三成但蒸馏损失的输入变少精度通常会掉 0.5 个点左右。如果显存实在不够另一个更划算的办法是减小局部裁剪的 batch 分组让局部视图只经过 key encoder不进入 query encoder这样显存消耗主要来自 key 分支而 key 分支不需要梯度可以开启torch.no_grad()并用 half 精度存储局部视图特征。这个改动不影响损失函数形态是工程实现里常见的优化策略也是「数据源码」这类标题下最容易抄的作业。本文还有配套的精品资源点击获取
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门