On-Policy无监督自蒸馏:无需标签与教师
先说结论这个方向解决的是“自蒸馏依赖标签或外部教师”的痛点试图在完全无监督的情况下让模型通过自身当前策略采样的数据完成蒸馏训练。标题里最关键的两个词是“On-Policy”和“without Any Supervision”。前者决定数据怎么采、loss 怎么算后者决定方法能不能脱离人工标注和预训练教师独立工作。如果你之前接触过 PPO应该对 On-Policy 不陌生——PPO 是强化学习里最典型的 On-Policy 算法它要求每次更新必须使用当前策略采样的轨迹。而这里把同样的思想搬到了自蒸馏里意味着蒸馏过程需要跟随模型自身的分布演化实时采样而不是在一个固定数据集上反复消费。这篇文章会把题目拆开讲清楚自蒸馏为什么值得做、无监督条件意味着什么、On-Policy 在这里到底起什么作用以及一套可落地的训练流程设计和验证思路。面向的是想理解自蒸馏原理、想复现类似方法、或者正在做无监督表示学习的读者。1. 核心概念速览先给一张速览表后面所有内容都围绕这几个点展开。能力项说明研究方向无监督自蒸馏Unsupervised Self-Distillation核心前提不使用人工标签、不依赖外部预训练教师模型方法特征On-Policy蒸馏样本来源于模型当前状态下的动态采样关键机制学生网络与教师网络协同教师提供学习目标学生从自身分布中持续学习与传统蒸馏区别传统蒸馏依赖固定教师和标注数据这里教师和学生同源且无需标注与 PPO 关系继承 On-Policy 思想每次更新前重新采样当前策略的数据分布适用任务表示学习、自监督预训练、视觉特征提取、强化学习状态表征等硬件门槛依赖具体模型规模自蒸馏通常比监督蒸馏更消耗训练数据与算力是否需要人工标注不需要是否支持批量训练支持且批量大小会直接影响分布估计质量适合读者算法工程师、研究生、自监督学习与强化学习入门者这个方向最大的价值不是发明了一种新 loss而是把“分布匹配”从静态变成了动态。传统自蒸馏如果完全用 ImageNet 或某个静态数据集那么数据分布是固定的模型学到的目标也是固定的。On-Policy 的做法是让模型一边学、一边改变自己遇到的样本分布这更接近强化学习里的 exploration 和 exploitation 平衡。2. 自蒸馏和它要解决的问题2.1 蒸馏是怎么工作的知识蒸馏是 Hinton 在 2015 年左右提出的经典训练范式。传统做法是训练一个大模型作为教师然后用教师输出的软标签去指导小模型训练。关键在于软标签比 hard label 包含更多类间关系信息例如“猫”和“狗”的输出分布比“猫”和“汽车”更接近学生能从中学到更平滑的决策边界。在自蒸馏里教师和学生不再是两个独立模型。教师通常由学生自身演化而来常见做法是使用学生参数的指数移动平均EMA或历史 checkpoint。这样做的收益是明显的不依赖额外的大模型推理也没有教师和学生能力差距过大的问题。2.2 为什么还需要“无监督”传统蒸馏有一个天然限制它需要标注数据来定义学习目标。哪怕教师模型在某个任务上学得很好如果换到新领域、新任务没有标签就无法从头开始蒸馏。而无监督自蒸馏要去掉这个限制。学习目标从“匹配教师输出的类别分布”变成“匹配特征空间中的一致性关系”例如同一图片的不同增强视角应该拥有相近的表征不同样本之间的关系结构应该保持稳定。这样就不需要任何标签甚至不需要知道类别数量。“without Any Supervision”强调的是整个训练闭环中不存在任何外部监督信号包括标签、伪标签、人工设计的度量标准都不是必须的。学习信号完全来自样本之间、视角之间、以及当前模型与历史自身之间的相对关系。2.3 难点在哪里无监督自蒸馏看起来美妙但实际实现有三个明显难点第一目标不稳定。没有标签等于没有一个固定锚点。模型在训练初期输出完全随机拿随机输出当目标很容易陷入 trivial solution比如所有样本都输出同一个向量。第二数据分布会漂移。随着模型更新模型对样本的“关注程度”也会变化。在静态数据集上这个问题还不明显如果是在线采样或强化学习环境中模型当前策略决定了下一次会遇到什么样本这就是 On-Policy 要处理的核心问题。第三特征空间需要额外约束。没有监督时特征表示可以任意旋转、缩放所以通常要加 uniformity 或 variance 之类的正则约束防止表征坍缩。3. On-Policy 与 PPO为什么它这么重要3.1 为什么说 PPO 是 On-Policy很多人第一次接触 On-Policy 都是从 PPO 开始的。PPO 全称是 Proximal Policy Optimization它在每次参数更新前必须使用当前策略与环境交互采集一批轨迹数据。更新完参数之后刚才采集的数据就不能再继续用于大规模训练因为策略变了数据分布已经过时。这就是 On-Policy 的本质训练数据必须来自当前状态下的行为策略。与之对应的是 Off-Policy典型代表是 DQN它可以把过去所有经验存进 replay buffer反复采样使用因为价值函数的更新不需要严格匹配当前策略。PPO 之所以经典是因为它用重要性采样和 clip 操作让样本利用率提升但又不至于完全脱离当前策略分布。它本质上是在 On-Policy 框架下尽可能延长样本使用时间而不是推翻 On-Policy。3.2 On-Policy Self-Distillation 中的“On-Policy”指什么在这个题目里On-Policy 的含义需要从两个层次理解。第一个层次是数据采样层次。如果训练数据来自一个不断变化的生成过程比如强化学习环境、模型自身生成的数据或者是每轮训练后重新从分布中采样那么当前模型只能看到当前策略产生的样本。这与 PPO 的环境交互逻辑一致。第二个层次是目标构造层次。自蒸馏需要教师网络提供目标而教师网络如果是学生网络的历史状态EMA那么教师更新的速度、学生采样的分布、以及目标分布之间必须保持同步。如果教师更新太快目标不稳定如果教师更新太慢目标又无法反映当前策略。这种同步性就是 On-Policy 思想的体现。需要澄清的是On-Policy 并不是说这里用了 PPO 算法而是说这个自蒸馏方法遵循了“当前策略采样、当前策略学习”的原则。3.3 On-Policy 和 Off-Policy 自蒸馏的对比对比维度On-Policy 自蒸馏Off-Policy 自蒸馏数据来源当前策略或当前模型分布下采样固定数据集 / replay buffer目标稳定性随策略变化需要 EMA 平滑相对稳定样本利用率低需要持续采样高可反复使用实现复杂度较高需要管理数据流较低适用场景在线学习、强化学习、动态环境静态数据集、离线预训练从这张表可以看到On-Policy 自蒸馏并不是在所有场景下都更优。如果你只是一个静态图片数据集完全离线训练那么 Off-Policy 方式是更务实的选择。On-Policy 的价值主要体现在模型会反过来改变数据分布的场景中。4. 方法设计拆解无监督自蒸馏的四个关键模块假设我们要从头实现一个 On-Policy 的无监督自蒸馏方法那么至少需要四个模块数据流构造、教师网络更新、学习目标构造、损失函数设计。4.1 数据流构造On-Policy 的数据流不能是一次性加载的静态数据集而应该是一个“采样器”。每一步训练采样器根据当前模型状态生成或选择一批样本。在强化学习场景中这个采样器就是环境交互器模型输出动作后环境返回下一个状态、奖励、终止信号。在自监督视觉场景中采样器是数据增强管道每次取出同一张图的不同增强视角。更一般的情况是模型自身作为生成器比如用当前模型生成文本、图像或伪样本然后拿这些生成结果作为蒸馏输入。从工程实现上看数据流构造是四个模块里最影响训练效率的一个。因为 On-Policy 要求每更新几步就要重新采样这会打断 GPU 的流水线并行所以通常需要独立的数据进程并把采样和训练解耦。4.2 教师网络更新教师网络是蒸馏目标的来源。在无监督自蒸馏里教师不能是外部预训练模型否则就违背了“without Any Supervision”的前提所以教师必须从学生演化而来。最常用的更新方式是 EMAθ_teacher m * θ_teacher (1 - m) * θ_studentm 是一个接近 1 的动量参数比如 0.99、0.999 或 0.9995。动量越大教师更新越慢目标越稳定但过大会导致教师跟不上学生蒸馏失去意义。实践中通常需要根据训练步数做 warmup或者使用 cosine schedule 动态调整 m。EMA 的一个额外好处是它天然提供了“历史平均”的特性。教师可以看成学生在不同训练时刻参数的加权平均这种平滑状态往往比当前状态的泛化能力更强。这也是 BYOL、DINO 等自监督方法能够稳定训练的原因之一。4.3 学习目标构造无监督蒸馏的核心难点在于没有标签怎么定义“学对了”。目前主流做法有以下几类第一视角一致性。对同一输入做两次不同的随机增强要求学生网络对这两个增强的输出在特征空间尽量一致。这是一种典型的自监督目标但单独使用容易导致表征坍缩。第二特征关系保持。不仅要求同一输入的两个视角一致还要求批次内不同样本之间的相似度关系一致。比如样本 A 和 B 的相似度应该大于 A 和 C 的相似度这个关系可以用对比损失约束。第三聚类一致性。将特征空间划分成若干簇要求同一簇内的样本具有相似特征且不同簇之间有区分度。因为没有标签簇的分配通常由教师网络生成学生则去匹配教师的簇分配结果。第四分布的 self-training。教师网络输出一个概率分布学生网络输出另一个分布两者做 KL 散度或交叉熵约束。这种方法需要避免模型对所有样本都输出均匀分布因此会加入锐化或温度调节。从“without Any Supervision”的角度来看以上目标都不需要标签但设计者仍然需要选择归纳偏置。比如选择视角一致性就隐含假设了“语义在增强下不变”这在部分领域不一定成立。从“without Any Supervision”的角度来看以上目标都不需要标签但设计者仍然需要选择归纳偏置。比如选择视角一致性就隐含假设了“语义在增强下不变”这在部分领域不一定成立。4.4 损失函数设计一种典型的无监督自蒸馏损失可以写成L λ_consistency * L_consistency λ_uniform * L_uniform λ_stability * L_stability其中L_consistency 是学生输出与教师输出的一致性损失 L_uniform 是特征均匀性约束防止坍缩 L_stability 是当前模型与历史状态之间的稳定性约束避免灾难性遗忘。具体用哪种距离度量取决于任务类型。如果输出是概率分布KL 散度更自然如果输出是特征向量负余弦相似度或 InfoNCE 是常用选择。5. 训练流程从零实现一个 On-Policy 自蒸馏框架下面给出一套通用的训练流程伪代码结构上适合理解实际项目需要按任务替换数据流和网络结构。import torch import torch.nn as nn import torch.nn.functional as F class StudentNet(nn.Module): def __init__(self, feature_dim128): super().__init__() self.encoder nn.Sequential( nn.Linear(784, 512), nn.ReLU(), nn.Linear(512, 256), nn.ReLU(), ) self.projection nn.Linear(256, feature_dim) def forward(self, x): h self.encoder(x) return self.projection(h) class TeacherNet(nn.Module): def __init__(self, student_net, momentum0.999): super().__init__() self.student_net student_net self.momentum momentum self._copy_parameters() def _copy_parameters(self): for teacher_param, student_param in zip(self.parameters(), self.student_net.parameters()): teacher_param.data.copy_(student_param.data) teacher_param.requires_grad False torch.no_grad() def update(self): for teacher_param, student_param in zip(self.parameters(), self.student_net.parameters()): teacher_param.data self.momentum * teacher_param.data (1 - self.momentum) * student_param.data def forward(self, x): return self.student_net(x)训练主循环def train_step(student, teacher, optimizer, x1, x2, temperature0.1): # x1, x2 是同一批样本的两个增强视角 student.train() # 学生输出 z1 student(x1) z2 student(x2) # 教师输出不计算梯度 with torch.no_grad(): t1 teacher(x1) t2 teacher(x2) # 一致性损失让学生的特征分布匹配教师 loss_consistency F.mse_loss(F.normalize(z1, dim-1), F.normalize(t2, dim-1)) loss_consistency F.mse_loss(F.normalize(z2, dim-1), F.normalize(t1, dim-1)) # 均匀性损失防止特征坍缩 z torch.cat([F.normalize(z1, dim-1), F.normalize(z2, dim-1)], dim0) loss_uniform torch.logsumexp(z z.T / temperature, dim1).mean() loss loss_consistency 0.01 * loss_uniform optimizer.zero_grad() loss.backward() optimizer.step() teacher.update() return loss.item()这里的核心逻辑是教师通过 EMA 从学生参数演化而来。学生需要同时拟合教师输出并在特征空间上保持分布均匀。每一步训练后教师更新然后下一步采样新数据。如果没有数据采样的动态变化这套流程就更接近 BYOL 或 DINO不能严格称为 On-Policy。要让它变成 On-Policy就需要把 x1、x2 的来源改为“当前策略产生的数据”而不是固定 DataLoader。6. On-Policy 数据采样的工程实现6.1 为什么 DataLoader 不够标准 PyTorch 训练用 DataLoader 加载固定数据集每个 epoch 只是对同一批数据做不同增强。这种数据流在分布上是静态的模型怎么更新都不影响下一批样本的来源。但 On-Policy 要求数据来源随模型变化。典型场景是强化学习中模型当前策略决定动作动作影响环境状态环境返回下一帧观测生成式自蒸馏中模型当前权重生成新的训练样本主动学习中模型当前置信度决定去哪些未标注数据。这种情况下每训练若干步就需要重新采样一批数据。数据采样过程本身也是计算过程所以必须把它从训练主线程中分离出来。6.2 一个简单的采样器接口class OnPolicySampler: def __init__(self, env, policy_net, num_steps1000): self.env env self.policy_net policy_net self.num_steps num_steps def sample_batch(self, batch_size256): observations [] for _ in range(batch_size): obs self.env.reset() for _ in range(self.num_steps): action self.policy_net(obs) obs, reward, done, info self.env.step(action) observations.append(obs) if done: break return torch.stack(observations)实际工程中采样器和训练器通常跑在不同的进程中采样器拿到当前模型参数产生一批数据传给训练器训练器更新参数后再把新参数同步给采样器。这个同步过程越频繁On-Policy 特性越严格但通信开销也越大。6.3 样本利用率与批次大小On-Policy 方法的一个实际问题是样本利用率低。Off-Policy 可以把一条样本反复用几十次On-Policy 往往只能用一次或者几次。因此On-Policy 自蒸馏通常需要更大的吞吐量来弥补样本效率损失。批次大小也会影响目标质量。如果批次太小batch 内样本之间的一致性估计方差大教师目标和学生目标都比较不稳定。在自蒸馏中千人以上规模的 batch 是常见配置但这对显存和算力提出了更高要求。如果设备有限可以先从 256 或 512 开始观察 loss 曲线是否平滑。7. 实验验证设计怎么看这个方法有没有效果单看 loss 下降是不够的因为无监督自蒸馏的 loss 可能因为坍缩而下降得很漂亮。需要设计一套完整的验证流程。7.1 训练阶段观察指标指标观察目的健康状态蒸馏 loss学生是否在拟合教师缓慢下降或震荡下降特征均匀性是否出现坍缩数值保持稳定且远离极端值教师-学生一致性EMA 平滑是否有效维持在一个合理范围内特征分布可视化判断特征的语义结构不同簇能自然分离如果蒸馏 loss 快速降到 0但特征分布可视化显示所有样本挤在一起说明模型已经坍缩loss 设计有问题。如果特征分布分散但没有任何语义结构说明目标的语义信号不够强。7.2 下游任务评估无监督学习的最终判断标准是下游任务表现。通常做法是把训练好的特征提取器固定只训练一个线性分类头观察分类准确率。这个做法叫 linear probing。在强化学习场景中可以把学到的状态表征输入策略网络比较使用固定随机特征、监督特征、无监督自蒸馏特征时的任务回报差异。7.3 消融实验为了验证 On-Policy 组件是否必要应该做以下三组消融去掉动态采样改为固定离线数据其它条件不变观察性能是否下降。去掉 EMA 教师让教师直接复制学生参数观察训练是否崩溃。去掉均匀性正则观察特征是否坍缩。这三组实验能分别证明 On-Policy 采样、目标平滑、坍缩约束各自的必要性。如果论文或项目中做了这些实验结论会比单纯给 loss 公式更有说服力。8. 训练资源需求与性能观察无监督自蒸馏的参数规模和显存占用没有一个固定数字取决于网络结构、batch size、输入分辨率、是否使用混合精度等。这里给出观察方法和优化建议而不是一个编造的显存数值。8.1 显存占用观察方法在 PyTorch 中可以用如下方式观察显存import torch print(fallocated: {torch.cuda.memory_allocated() / 1024 ** 3:.2f} GB) print(freserved: {torch.cuda.memory_reserved() / 1024 ** 3:.2f} GB)自蒸馏通常比普通分类任务更消耗显存因为需要同时计算学生和教师两个分支教师分支虽然不计算梯度但仍然需要前向传播占用额外显存。如果显存不足可以按以下顺序优化降低 batch size但要注意 batch 太大会影响分布估计使用混合精度训练显存占用可降低约 40%教师分支使用更小的输入分辨率学生分支保持正常分辨率使用 gradient checkpointing 降低激活值显存把教师前向传播放到另一个 GPU 或使用 CPU 前向。8.2 CPU 与 GPU 训练差异自蒸馏的教师-学生双分支结构在前向传播上的计算量约为普通单模型的 2 倍。CPU 训练在数据规模较小时可以完成功能验证但 On-Policy 需要高频采样CPU 很难满足吞吐要求。一个务实做法是先用 CPU 在小规模数据上验证 loss 能正常下降、教师更新逻辑没有 bug再切换到 GPU 跑完整实验。这样能避免把调试时间浪费在等待大模型训练上。8.3 训练稳定性观察On-Policy 自蒸馏最大的敌人是训练不稳定。教师更新太慢导致目标滞后学生可能在一个已经过时的分布上过拟合教师更新太快导致目标震荡学生可能永远追不上教师。常见的稳定性指标是教师网络输出的变化幅度。可以周期性地计算教师输出分布的均值漂移如果漂移过大就需要增大动量参数 m。9. 常见问题与排查方法问题现象可能原因排查方式解决方案训练 loss 快速下降但下游任务表现差特征坍缩可视化特征分布增加均匀性正则调整温度参数教师更新跟不上学生EMA 动量太大观察教师输出漂移幅度减小动量或增加 Teacher warmup学生与教师输出差异过大学生初始化或学习率过高打印梯度统计降低学习率使用更平滑的优化器On-Policy 采样导致训练吞吐下降采样和训练串行查看 GPU 利用率使用多进程采样做预取缓冲同一批数据反复训练但效果变差样本被过度消费观察每个 epoch 的 loss减少重复次数增加新采样频率特征均匀但无语义目标信号弱做 linear probing增强数据增强多样性或改用对比损失显存溢出batch 太大或双分支显存占用高逐步减小 batch 测试混合精度、梯度检查点、教师低分辨率其中特征坍缩是最隐蔽的问题因为它不会让训练报错也不会让 loss 明显异常但会让整个模型失去意义。任何一次实验都建议先跑一个很小规模的可视化验证确认特征空间真的有结构再继续大规模训练。10. 最佳实践与使用建议从工程角度看这个方向值得投入但需要控制实验预期。无监督自蒸馏不是一键训练就出效果的算法它在数据规模、批次大小、增强策略、动量调度上都有较多敏感点。建议按照以下顺序推进实验第一先跑通最小可运行版本。用一个简单数据集如 MNIST 或 CIFAR-10 的子集验证学生、教师、EMA 更新、loss 计算整个链路没有 bug。这一步不要追求精度只追求不崩溃、不坍缩。第二固定一套默认超参数。把 momentum、temperature、batch size、学习率记录为一个配置文件后续调参时只改一个变量不要同时修改多个。第三验证 On-Policy 的价值。找一个动态数据场景比如环境交互或生成式采样对比固定数据集和动态采样的结果。如果两者效果接近说明在当前任务上 On-Policy 不是关键因素不必增加工程复杂度。第四保存所有 checkpoint 和特征可视化结果。无监督训练过程中模型质量不是单调上升的可能在中途达到最好效果。没有 checkpoint 就丢失了最佳模型。第五注意合规边界。如果实验数据涉及人物肖像、版权音频、敏感文本必须确认数据来源合法并取得授权。自监督方法虽然在训练中不使用标签但它会学习数据的内在分布如果数据本身存在偏见或不合法内容学到的特征也会受到影响。商用前必须对训练数据做审查。11. 扩展方向与后续思路这个方向可以延伸开来的点很多。从场景上可以尝试把 On-Policy 自蒸馏与强化学习结合让智能体在探索过程中积累的状态直接作为自蒸馏数据既做表征学习又做策略优化。从结构上可以用 Vision Transformer 替代简单 MLP观察自蒸馏在更大模型上的表现。从目标设计上可以尝试在特征空间引入结构化约束比如把特征分组、加入可解释维度。如果你正在做自监督学习相关研究这个方向值得关注的核心不是“无监督”这个标签而是 On-Policy 带来的数据分布动态性。它把蒸馏从一种静态匹配方法变成了一个与模型策略共同演化的过程。这个演化过程是否稳定、是否能带来更高质量的特征是未来可以深入验证的问题。最值得先做的一步是把本文第 5 节的伪代码跑通然后换成你的真实数据和网络结构观察特征可视化结果。如果特征有结构再逐步引入 On-Policy 数据采样如果特征坍缩优先调整均匀性约束和温度参数。整个方向的关键在于稳定性和数据的动态匹配而不是追求某一次 loss 数字的下降。