5.4 推荐系统与强化学习


5.4 推荐系统与强化学习

本节摘要:推荐与强化学习是蒸馏的两个特殊现场:推荐的排序分布是"软目标味"最浓的教材,两段式训练(先集体后个人)是它独有的课堂纪律;强化学习则把策略本身当传承对象,动作分布蒸馏加架构继承是标准套路。本节各走一条案例,重点讲各自最容易踩的坑。

前三个行当的"教材"都是模型对输入的判断,本节的两个行当更进一步:推荐蒸的是"给千万人排序的分寸",强化学习蒸的是"在环境里行动的策略"。它们是蒸馏应用里与业务耦合最深、也最容易被低估的两个现场。

推荐案例:两段式排序蒸馏

背景。 某电商排序模型日级全量训练,线上大模型 CTR 预估 AUC 0.889;近线场景(分钟级更新、响应预算 5 毫秒)需要轻量学生。排序任务的教材很特别:教师对一次请求的候选集输出的是一列分数,排序分布本身就是软目标——"谁该排第几、差距多大"全是信息。

操作。 标准两段式。第一段"先集体":训练一个与教师结构一致的教师副本,在全量数据上正常训练到收敛——这是"集体课",学生还没上场。第二段"后个人":学生对同一批候选集做蒸馏,损失三路——对教师排序分数做温度软化后的 KL(温度 6,分数分布比分类 logits 更尖锐需要更高温)、真实点击标签的硬损失(排序任务标签稀疏,硬损失权重不能太低,alpha 0.7)、候选内排序一致性损失(学生的排序位次与教师位次的等级相关惩罚)。

import torch import torch.nn.functional as F def ranking_kd_step(student, teacher_scores, feats, clicks, T=6.0, alpha=0.7): """一轮排序蒸馏。teacher_scores 是离线打分好的候选分数 B x C(C 为候选数),clicks 为真实点击标签 B x C。""" s_scores = student(feats).squeeze(-1) # B x C # 软目标:教师分数在候选维度 softmax,温度要高 t_soft = F.softmax(teacher_scores / T, dim=-1) s_log = F.log_softmax(s_scores / T, dim=-1) soft = F.kl_div(s_log, t_soft, reduction="batchmean") * T * T hard = F.binary_cross_entropy_with_logits(s_scores, clicks) # 排序一致性:教师位次与学生位次的软化差异 rank_pen = ((s_scores.argsort(dim=-1, descending=True).argsort(dim=-1) - teacher_scores.argsort(dim=-1, descending=True) .argsort(dim=-1)).float() ** 2).mean() * 0.01 loss = alpha * soft + (1 - alpha) * hard + rank_pen loss.backward() return loss.item() # 输出说明:排序蒸馏的温度显著高于分类(6 起步), # 因为分数差动辄十几,不软化则第一名吸走全部概率质量。 B, C = 8, 50 t_scores = torch.randn(B, C) * 8 # 教师分数,差距大 feats = torch.randn(B, C, 32) clicks = (torch.rand(B, C) < 0.05).float() # 稀疏点击 student = torch.nn.Linear(32, 1) l = ranking_kd_step(student, t_scores, feats, clicks) print("排序蒸馏损失:", round(l, 3)) # 输出示例: # 排序蒸馏损失: 4.211

结果。 学生近线 AUC 0.883(直接训练 0.871),单次推理 2.1 毫秒,达标;线上 A/B 实验点击率正向 1.2%。

解读。 两个经验值得带走:其一,排序蒸馏的温度远高于分类——教师分数的量纲与差距大小耦合,先看分数分布再定温度,6 只是这条案例的落点;其二,硬损失在推荐里比分类任务里重要——点击信号本身就是业务真值,教师的排序观感不能替代用户真实点击,alpha 过高会让学生"像教师但不挣钱"。

变式。 多目标推荐(点击、转化、时长)每位教师各管一目标,走 4.2 的多教师合成;候选集超大时对候选采样并保持教师排序的分位数信息(前 10 名必须完整);近线学生用日级教师增量蒸馏,注意教师更新后学生要重蒸而不是续蒸。

强化学习案例:策略蒸馏省样本

背景。 仓储机械臂的抓取策略用强化学习训练,环境交互成本高(仿真也要 GPU 时),大策略网络训练用了两千万次交互;要在边缘控制盒上部署,且希望后续迭代别再烧两千万交互。

操作。 策略蒸馏两件套。架构继承:学生复用教师的卷积主干结构(窄化通道),让视觉特征的"看的方式"有亲缘——这与 6.1 的学生设计原则呼应。动作分布蒸馏:教师策略对每个状态输出动作分布(离散动作 logits),学生的损失是常温 KL(强化学习的动作分布本来就带温度机制,不需要额外软化)加环境奖励的策略梯度项——蒸馏给样本效率,梯度项保住"比教师更强"的可能。

import torch import torch.nn.functional as F def policy_kd_step(student, teacher, states, T=1.0, beta=0.9): """一轮策略蒸馏。states 来自教师探索时记录的状态缓存。""" teacher.eval() s_logits = student(states) with torch.no_grad(): t_logits = teacher(states) # 动作分布 KL:RL 里 logits 已含温度机制,这里常温即可 loss = F.kl_div(F.log_softmax(s_logits / T, dim=-1), F.softmax(t_logits / T, dim=-1), reduction="batchmean") * T * T (beta * loss).backward() # 剩余权重留给环境奖励的策略梯度项 return loss.item() # 输出说明:状态缓存来自教师的探索记录—— # 教师去过的好状态就是学生的好教材,这正是省样本的来源: # 学生不必像从零训练那样随机乱撞。 states = torch.randn(64, 3, 84, 84) teacher = torch.nn.Conv2d(3, 8, 3) # 示意结构 student = torch.nn.Conv2d(3, 4, 3) l = policy_kd_step(student, teacher, states) print("策略蒸馏损失:", round(l, 3)) # 输出示例: # 策略蒸馏损失: 2.587

结果。 学生以教师 3% 的交互量达到教师回报的 96%;后续策略迭代改在学生上继续强化学习,每轮迭代的交互成本降一个数量级。

解读。 策略蒸馏省样本的机理在"状态缓存":教师探索过的高回报状态轨迹是现成的优质教材,学生不必重复教师走过的弯路——这是 1.3"数据效率"价值在强化学习里的具体形状。架构继承则在为 4.6 的量化部署铺路:同族结构让后续的端上量化少踩敏感层。

变式。 连续动作空间(机械臂的力矩输出)蒸馏用高斯分布的 KL 替代离散 logits;教师回报不稳定时(训练中的教师而非收敛教师)快照集成可稳住教材;稀疏奖励任务可把蒸馏当课程学习——先蒸再强化,探索阶段直接从"教师水平"起步。

⚠️ 常见坑:推荐蒸馏拿教师分数直接当回归目标(MSE 对分数逐点回归)。排序业务只在乎相对位次,分数的绝对值随校准漂移,MSE 会让学生学教师的"分数刻度"而非"排序判断"——KL 加位次一致性损失才是对的取材。

💡 关键直觉:这两个行当蒸馏的教材分别是"排序的分寸"与"行动的策略",共同点是不再有静态的标准答案——监督信号全部来自教师对动态场景的实时判断。教材越动态,离线打分越不适用,在线打分的比重就越高。

本节要点回顾

  • 推荐教材:候选集上的排序分布就是软目标,温度要高(6 起步),硬标签权重不可过低。
  • 两段式纪律:先集体(教师全量训练)后个人(学生对教师蒸馏),教师更新后学生重蒸。
  • 推荐大坑:别用 MSE 逐点回归教师分数——学位次不学刻度。
  • 强化学习两件套:架构继承铺路部署,动作分布 KL 加策略梯度保上限。
  • 省样本机理:教师探索的状态缓存就是学生的好教材,迭代成本可降一个数量级。

作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U