在搜索、推荐和问答系统的精排阶段,重排序模型经常选择交互式结构。查询和文档被拼接成一对输入,让Transformer同时看到两侧token,这样的cross-encoder精度高,但单次前向成本也高。一旦请求里包含几百条候选,线上延迟会迅速上升。解决这个问题不能只靠堆GPU,更实际的做法是减少模型结构开销和数值计算开销,也就是蒸馏和量化部署。

一、先定位重排序为什么慢
交互式模型的计算方式决定了它的复杂度。对于每一个query-doc对,模型都要重新计算query token的self-attention。即使query完全相同,只要文档变化,整条序列的注意力矩阵都会改变。因此无法像双塔召回那样把文档向量离线算好,线上只算点积。候选数量N、序列长度L和层数H叠加起来,延迟近似按N乘H乘L的平方增长。
举个例子,假设精排候选为200条,每条拼接后截断到128 token,使用12层BERT结构的cross-encoder。在单卡推理时,只跑模型部分可能就需要几十毫秒到上百毫秒,峰值并发下P99延迟会更高。再加上tokenize、padding、batch调度和结果排序等外围开销,重排阶段容易成为整个链路中最慢的一环。双塔模型虽然快,但查询和文档在最后才做点积,缺少细粒度交互,相关性通常会弱一些。
除了模型结构,推理框架也会带来额外成本。PyTorch动态图在小batch下调度开销明显,一些算子没有融合,也会浪费计算资源。因此优化时要同时看三块:模型结构能否简化、数值精度能否降低、推理图能否导出并融合。
二、蒸馏:把教师模型的能力迁移到轻量结构
蒸馏不是简单地把模型层数砍掉,而是让学生模型学习教师模型输出的分数分布。教师通常是效果最好的大号cross-encoder,学生可以选择两种结构。第一种是同构学生,仍然保留查询和文档的交互编码方式,但减少层数或隐层维度。第二种是异构学生,例如改成双塔,查询和文档分别编码后计算相似度。前者延迟降幅有限但相关性损失较小,后者速度接近召回但交互能力会下降。
同构蒸馏训练时,教师模型对每个query-doc对打一个相关性分数,学生模型也输出分数,两者通过KL散度对齐。同时还要加入真实点击或标注样本的交叉熵损失,让学生不只模仿教师,也能和实际反馈保持一致。温度参数用来控制软标签的平滑程度,温度越高教师输出越平均,学生能学到更多类间关系,但也要防止过度平滑。
import torch
import torch.nn as nn
import torch.nn.functional as F
class DistillLoss(nn.Module):
def __init__(self, temperature=3.0, alpha=0.5):
super().__init__()
self.temperature = temperature
self.alpha = alpha
def forward(self, student_logits, teacher_logits, labels):
soft_teacher = F.log_softmax(teacher_logits / self.temperature, dim=-1)
soft_student = F.log_softmax(student_logits / self.temperature, dim=-1)
distill_loss = F.kl_div(soft_student, soft_teacher, reduction='batchmean')
ce_loss = F.cross_entropy(student_logits, labels)
return self.alpha * distill_loss + (1 - self.alpha) * ce_loss
# 在训练循环中先固定教师模型
teacher_model.eval()
with torch.no_grad():
teacher_score = teacher_model(input_ids, attention_mask, token_type_ids)
student_score = student_model(input_ids, attention_mask, token_type_ids)
loss = distill_loss(student_score, teacher_score, labels)
loss.backward()
如果要把cross-encoder蒸馏到双塔,损失函数基本相同,只是学生输出的logits改为query向量和doc向量的余弦相似度。线上部署时,文档侧向量可以离线批量生成并存入向量库,请求阶段仅编码query向量,再与候选文档向量做近邻计算。这种方式能把200条候选的Transformer前向压缩到1次query编码加一次矩阵乘法,延迟下降非常明显。
蒸馏过程对训练数据要求较高。只用随机负样本,学生模型容易只学会区分明显不相关的文档,对困难样本排序能力不足。可以从线上日志里挖掘曝光未点击或点击排名靠后的文档作为难负样本,也可以让教师模型对大规模候选打分,挑出分数接近的pair让学生重点学习。温度一般取2到5,学生结构越小,往往需要更多训练步数和更大的batch来稳定收敛。
三、量化:从FP32到FP16和INT8
量化降低的是模型前向的数值精度和存储成本。FP32权重每个参数占4字节,FP16占2字节,INT8只占1字节。显存下降之外,整数计算还可以利用现代CPU和GPU的向量化指令,提高吞吐。对重排序这种小batch、高延迟敏感场景,FP16通常风险较低,INT8则需要更谨慎地评估相关性损失。
PyTorch对Transformer模型提供了动态量化接口,主要把Linear层权重转为qint8,激活在推理时根据输入范围动态计算。这种方式不需要额外提供校准数据集,落地成本低。下面是一个简化的示例,实际使用时需要替换成自己的重排模型结构。
import torch
from torch.quantization import quantize_dynamic
class SimpleRanker(torch.nn.Module):
def __init__(self):
super().__init__()
self.proj = torch.nn.Linear(768, 256)
self.out = torch.nn.Linear(256, 1)
def forward(self, x):
return self.out(torch.relu(self.proj(x)))
model = SimpleRanker()
# 动态量化Linear层,权重转成int8
quantized_model = quantize_dynamic(model, {torch.nn.Linear}, dtype=torch.qint8)
with torch.no_grad():
dummy = torch.randn(8, 768)
fp_output = model(dummy)
q_output = quantized_model(dummy)
print(torch.max(torch.abs(fp_output - q_output)))
如果希望激活也量化成INT8,可以走静态量化或量化感知训练。静态量化需要准备代表性数据,统计每层激活的数值范围,把激活也映射到8位整数。量化感知训练则在训练阶段插入伪量化节点,让模型参数提前适应量化误差,通常掉点更少。重排模型对排序精度比较敏感,INT8后如果NDCG下降超过千分之几,建议优先尝试QAT。
工程部署时可以把学生模型导出为ONNX,再用ONNX Runtime或TensorRT做图优化与量化。ONNX Runtime提供动态量化接口,TensorRT则能在GPU上完成INT8校准和算子融合。即使不做量化,导出ONNX后去掉Python动态图开销,推理延迟也可能有所下降。
from onnxruntime.quantization import quantize_dynamic, QuantType model_input = "reranker.onnx" model_output = "reranker_int8.onnx" # 动态量化导出int8权重 quantize_dynamic(model_input, model_output, weight_type=QuantType.QInt8)
四、组合策略与压测验证
蒸馏和量化通常组合使用效果更好。先通过蒸馏降低结构复杂度,把小模型的延迟拉下来,再在这个基础上做量化,进一步压缩权重和计算。如果顺序反过来,先量化大模型,延迟大头仍然来自结构本身,收益会比较有限。对效果要求高的场景,可以选择浅层cross-encoder学生加FP16;对延迟要求极高的场景,可以蒸馏成双塔再做INT8量化。
评估时不能只看模型单次前向耗时。要按线上真实候选数量压测,观察P50、P99延迟、QPS和显存占用。效果侧同步统计NDCG、MRR或点击率。下面是一组示意数据,用于帮助团队建立对比基线。
| 方案 | 候选200条P99延迟 | NDCG@10 | 说明 |
|---|---|---|---|
| 12层cross-encoder FP32 | 260ms | 0.832 | 基线 |
| 6层学生 FP32 | 140ms | 0.821 | 结构蒸馏 |
| 6层学生 INT8 | 75ms | 0.817 | 量化后 |
| 双塔学生 FP32 | 18ms | 0.795 | 牺牲交互 |
灰度发布时可以先切10%流量,观察排序结果和用户行为指标。如果延迟下降但点击率明显变差,通常是学生模型对难负样本的区分能力不足,需要补充难负样本重新蒸馏,或者提高蒸馏损失中教师软标签的权重。如果指标稳定,再逐步放量。对高价值流量和长尾流量也可以分别部署不同规格模型,让整体成本与体验达到平衡。
重排序延迟优化是结构选择、训练策略和推理工程共同作用的结果。蒸馏决定模型算多少,量化决定每次算多贵,推理框架决定实际执行效率。把这三件事分开验证、合并上线,才能在不明显牺牲排序质量的前提下解决线上慢的问题。