团队刚接了一个短视频推荐的项目,粗排层用的是一个超大的双塔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 注意事项

  1. 必须保证Teacher和Student的特征工程完全一致:包括特征的取值范围、归一化方式、特征的计算逻辑,哪怕改一个特征的阈值,都会导致分布错位;
  2. 蒸馏时一定要加入分布对齐损失:尤其是当特征维度高、特征的重要性差异大的时候,对齐分布能有效减少排序偏差;
  3. 线上测试必须测排序重合度:不能只看AUC或点击率,要测Top10的item重合度,这才是粗排的核心指标;
  4. 定期巡检特征分布:线上的用户行为会变化,Teacher的特征分布也会漂移,需要每月对比Teacher和Student的特征统计,及时重新蒸馏;

六、总结

这次踩的坑让我们明白,模型蒸馏不是“复制粘贴”大模型的能力,而是“学习大模型的逻辑和特征偏好”。之前我们只关注蒸馏的软标签,忽略了特征分布这个核心细节,导致线上排序完全跑偏。通过对齐Teacher和Student的特征分布,加入分布对齐损失,我们解决了线上一致性的问题,同时保留了模型压缩的收益。这个教训也给后续的模型迭代提了个醒:不管是做模型压缩还是小模型迭代,都要盯着“特征分布”这个底层逻辑,不然再小心都会出问题。