团队刚接了一个短视频推荐的项目,粗排层用的是一个超大的双塔Teacher模型,推理一次要花10毫秒,线上QPS一到10万就扛不住——毕竟短视频的请求是秒级的,峰值流量经常突破20万。为了给线上“减负”,我们决定用知识蒸馏做个Student小模型,目标是把Teacher的排序能力“压缩”到只有原来1/5大,推理速度提5倍。结果上线后麻烦来了:线上一致性校验直接红了,同一用户同一时间的请求,Student返回的粗排列表和Teacher的重合度只有62%,点击率比之前掉了3%左右。
一、问题触发:粗排双塔蒸馏后的线上异常
我们先把问题拆得通俗点:粗排的作用是从召回的1000个候选里挑100个给精排,相当于从全校1000个学生里挑前100名重点培养。Teacher是“金牌教练”,有多年带竞赛生的经验,Student是“新手教练”,学教练的思路来挑学生。本来以为新手教练能学会金牌教练的眼光,结果挑出来的学生完全不对路——金牌教练挑的是思维活跃的竞赛苗子,新手教练挑的是刷题多的高分选手,本质是看的维度不一样。
1.1 问题表现的具体现象
我们拉了A/B实验的明细数据,发现两个核心异常:一是线下测试时,Student模型对item的排序NDCG@10比Teacher低了8个点,说明排序逻辑完全跑偏;二是线上实时抓了1000个用户的请求,同一请求的Top10粗排结果,重合度仅62%,其中有38%的item是Teacher里排100名外的“边角料”,反而被Student放到了前面。最关键的是,之前我们做过多次模型小版本迭代,排序一致性都是90%以上,这次蒸馏后直接掉了30个点,明显不是模型容量的问题。
二、先搞懂:双塔蒸馏的核心逻辑
很多人对知识蒸馏的理解停留在“小模型学大模型的输出”,其实更准确的说法是“小模型学大模型的判断逻辑,尤其是特征分布的规律”。这里举个生活化的例子:Teacher(金牌教练)看一个学生,会综合看“数学兴趣、动手能力、过往竞赛成绩”三个维度,给每个学生打一个“竞赛潜力分”,比如A是9分(数学好+动手强+拿过奖),B是7分(数学好但动手弱),C是2分(数学差)。Student要学的不是A的9分具体是多少,而是这三个学生的潜力排序:A>B>C,这个排序背后藏着Teacher的特征分布——数学兴趣占60%权重,动手能力占30%,过往成绩占10%。如果Student学的时候,把动手能力的权重搞成50%,那他就会把B排到前面,就出现了错位。
2.1 为什么特征分布是核心?
双塔模型的本质是把用户和item都映射到同一个向量空间,然后算用户向量和item向量的相似度,排序的时候按相似度从高到低排。Teacher和Student的特征分布错位,就是指它们对“哪些特征更重要”的判断不一样——比如Teacher看用户最近1小时的点击,篮球占60%,足球占30%;但Student训练时用的样本里,足球占了50%,篮球占了30%,那Student在算相似度的时候,就会更看重足球相关的item,自然排序就偏了。
三、排查过程:定位到Teacher和Student的特征分布错位
一开始我们走了很多弯路:先是以为蒸馏的损失函数没调好,换了3种损失函数(L2损失、交叉熵、KL散度),问题依然存在;又以为是线上环境的特征工程和线下不一致,线下用了特征标准化,线上用了归一化,改完后重合度只涨了2个点;后来又换了更大的Student模型(参数是原来的2倍),结果还是不对。直到我们去查Teacher和Student的特征统计分布,才发现问题所在。
3.1 特征分布错位的具体表现
我们抽了100万条训练样本,统计Teacher和Student对“用户最近1小时点击类别”这个特征的分布:Teacher的分布是篮球61%、足球28%、其他11%;Student训练用的样本里,篮球只有38%,足球却有52%,其他10%——这说明Student的训练数据采样偏了,把足球的样本多采了,导致它学的特征分布和Teacher完全不一样,相当于新手教练的训练样本里,足球相关的优秀选手更多,自然挑出来的人就不对。
四、解决思路:对齐Teacher和Student的特征分布
找到问题后,我们想到两个解决方向:一是对齐两者的训练数据分布,二是在蒸馏时加入分布对齐的损失。我们用Python实现了第二种方式,因为成本更低,而且能直接加到原有蒸馏框架里。这里统一用Python技术栈,代码及注释如下: 技术栈:Python 3.8 + PyTorch 1.12
import torch
import torch.nn as nn
import torch.nn.functional as F
# 简化版双塔Teacher模型(实际项目中是大模型,这里用小模型模拟)
class TeacherModel(nn.Module):
def __init__(self, feat_dim=128):
super().__init__()
# 模拟Teacher的特征编码层,权重大,表达能力强
self.encoder = nn.Sequential(
nn.Linear(feat_dim, 256),
nn.ReLU(),
nn.Linear(256, feat_dim)
)
self.fc = nn.Linear(feat_dim, 1)
def forward(self, user_feat, item_feat):
# 双塔的核心:用户和item做内积后输出分数
user_emb = self.encoder(user_feat)
item_emb = self.encoder(item_feat)
return self.fc(torch.sum(user_emb * item_emb, dim=1, keepdim=True))
# 简化版双塔Student模型(线上部署的小模型)
class StudentModel(nn.Module):
def __init__(self, feat_dim=128):
super().__init__()
# Student的编码层更轻,参数少,推理快
self.encoder = nn.Sequential(
nn.Linear(feat_dim, 64),
nn.ReLU(),
nn.Linear(64, feat_dim)
)
self.fc = nn.Linear(feat_dim, 1)
def forward(self, user_feat, item_feat):
user_emb = self.encoder(user_feat)
item_emb = self.encoder(item_feat)
return self.fc(torch.sum(user_emb * item_emb, dim=1, keepdim=True))
# 核心:带分布对齐的蒸馏损失函数
def distill_with_align_loss(teacher, student, user_feat, item_feat, label, alpha=0.5):
# 1. 拿到Teacher的软标签和特征分布(不更新Teacher权重)
with torch.no_grad():
teacher_logits = teacher(user_feat, item_feat)
# 计算Teacher的特征统计分布:每个维度的平均激活值,代表Teacher的特征偏好
teacher_feat_dist = torch.mean(torch.abs(teacher_logits), dim=0)
# 2. 拿到Student的输出和特征分布
student_logits = student(user_feat, item_feat)
student_feat_dist = torch.mean(torch.abs(student_logits), dim=0)
# 3. 标准蒸馏损失:Student匹配Teacher的软标签
hard_label_loss = F.mse_loss(student_logits, label)
soft_label_loss = F.kl_div(
F.log_softmax(student_logits / 1.0, dim=0),
F.softmax(teacher_logits / 1.0, dim=0),
reduction='batchmean'
)
distill_loss = hard_label_loss + soft_label_loss
# 4. 新增:特征分布对齐损失,惩罚Student和Teacher的特征偏好差异
align_loss = F.mse_loss(student_feat_dist, teacher_feat_dist)
# 5. 总损失:蒸馏损失 + 分布对齐损失,权重alpha可以根据调参结果改
total_loss = distill_loss + alpha * align_loss
return total_loss
4.1 解决后的效果
加了分布对齐损失后,我们重新训练Student模型,上线后再测:线上粗排结果和Teacher的重合度涨到了91%,点击率回升了2.8%,和之前的表现基本一致,QPS也达到了原来的4.8倍,完美满足线上需求。
五、应用场景、技术优缺点及注意事项
5.1 应用场景
这个方案的核心是“模型压缩+线上一致性”,适合的场景是:线上粗排层需要极高的推理速度(QPS要求高),同时不能牺牲太多排序效果;线下有成熟的大模型Teacher,适合用来做知识蒸馏;召回层如果是全量候选,不需要用蒸馏(因为召回是匹配,不是排序),但粗排是精选,非常适合用这个方案。
5.2 技术优缺点
优点:Student模型小,推理速度快,能扛更高的QPS;保留了Teacher的排序能力,线上一致性好;比普通蒸馏多了分布对齐,解决了核心的错位问题; 缺点:需要额外加入分布对齐的损失,训练的时候要多一个步骤,调参成本略高;如果Teacher本身的特征分布在变化(比如用户兴趣随时间漂移),Student需要定期重新蒸馏,否则会再次出现错位;
5.3 注意事项
- 必须保证Teacher和Student的特征工程完全一致:包括特征的取值范围、归一化方式、特征的计算逻辑,哪怕改一个特征的阈值,都会导致分布错位;
- 蒸馏时一定要加入分布对齐损失:尤其是当特征维度高、特征的重要性差异大的时候,对齐分布能有效减少排序偏差;
- 线上测试必须测排序重合度:不能只看AUC或点击率,要测Top10的item重合度,这才是粗排的核心指标;
- 定期巡检特征分布:线上的用户行为会变化,Teacher的特征分布也会漂移,需要每月对比Teacher和Student的特征统计,及时重新蒸馏;
六、总结
这次踩的坑让我们明白,模型蒸馏不是“复制粘贴”大模型的能力,而是“学习大模型的逻辑和特征偏好”。之前我们只关注蒸馏的软标签,忽略了特征分布这个核心细节,导致线上排序完全跑偏。通过对齐Teacher和Student的特征分布,加入分布对齐损失,我们解决了线上一致性的问题,同时保留了模型压缩的收益。这个教训也给后续的模型迭代提了个醒:不管是做模型压缩还是小模型迭代,都要盯着“特征分布”这个底层逻辑,不然再小心都会出问题。
Comments