LinkedIn 公布了其 AI 求职搜索训练基础设施的一部分:通过多教师知识蒸馏,将多个大型教师模型中的排序能力压缩到一个仅有 0.6B 参数的紧凑模型中。这个思路的关键不只是“把大模型变小”,而是让多个教师模型共同提供训练信号,再用一个适合线上部署的学生模型完成最终排序。
对于搜索、推荐和广告系统来说,这种方案很有吸引力:大模型可以承担复杂的语义理解和特征组合,小模型则负责高吞吐、低成本的在线候选排序。
为什么单一教师模型不一定够用
传统知识蒸馏通常采用一个教师模型和一个学生模型:教师输出 soft label,学生学习教师的预测分布,同时结合真实标签训练。
但求职搜索的相关性并不只有一个维度。例如:
- 职位名称与查询词是否匹配;
- 候选人的技能、经验与职位要求是否匹配;
- 地点、远程办公、职级和行业是否匹配;
- 用户历史行为是否暗示了更强的兴趣;
- 结果是否满足业务排序约束。
不同教师模型可能在不同类型的信号上更擅长。一个教师模型擅长文本语义,另一个模型可能更擅长用户行为或结构化特征。多教师蒸馏的目标,就是把这些互补能力组合成更丰富的监督信号,而不是让学生模型只模仿某一个模型的偏好。
可以把训练过程抽象成下面的形式:
查询、用户、职位特征
│
├── 教师模型 A:语义相关性
├── 教师模型 B:行为相关性
└── 教师模型 C:结构化匹配
│
教师输出的分数或概率分布
│
加权、校准或融合后的蒸馏目标
│
0.6B 学生排序模型
这里的学生模型不需要复现教师的全部能力。它只需要在目标排序任务上保留最有价值的决策边界,并且足够小、足够稳定,适合大规模服务。
多教师蒸馏的训练目标
一个实用的多教师目标通常包含三部分:真实标签损失、教师软标签损失,以及可能存在的排序损失。
设真实标签为 y,第 i 个教师的输出为 p_i,学生输出为 p_s,则可以使用类似下面的目标:
L = α * L_hard(y, p_s)
+ β * Σ w_i * KL(p_i || p_s)
+ γ * L_rank
其中:
L_hard让学生模型继续遵守真实点击、申请或相关性标签;KL让学生学习教师输出中的“软信息”,例如样本之间的相对置信度;w_i表示不同教师的权重;L_rank可以用于强化 pairwise 或 listwise 排序目标;α、β、γ控制不同信号的影响。
温度参数也很重要。教师输出经过较高温度 T 的 softmax 后,类别之间的差异会变得更平滑,学生可以获得更多“哪些候选人略好、哪些明显不相关”的信息。工程上还需要注意蒸馏损失的尺度,常见做法是对 KL 损失乘以 T²,避免温度变化导致梯度规模失控。
一个可以运行和改造的 PyTorch 示例
下面的示例使用三个教师网络和一个学生网络,在合成的二分类排序数据上进行多教师蒸馏。它不是 LinkedIn 的生产实现,而是一个最小可运行模型,用来说明教师融合、硬标签和软标签如何组合。运行前安装 PyTorch:
python -m pip install torch
python multi_teacher_distill.py
将下面内容保存为 multi_teacher_distill.py:
import torch
import torch.nn as nn
import torch.nn.functional as F
torch.manual_seed(7)
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
class Ranker(nn.Module):
def __init__(self, input_dim: int, hidden_dim: int):
super().__init__()
self.net = nn.Sequential(
nn.Linear(input_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, 2),
)
def forward(self, x):
return self.net(x)
# 假设每个样本包含查询、用户和职位拼接后的特征。
num_samples, input_dim = 2048, 32
x = torch.randn(num_samples, input_dim, device=DEVICE)
true_w = torch.randn(input_dim, 1, device=DEVICE)
y = (x @ true_w + 0.3 * torch.randn(num_samples, 1, device=DEVICE) > 0).long().squeeze(1)
# 生产环境中,教师通常来自已经训练好的不同模型;这里用不同随机种子模拟互补教师。
teachers = []
for seed in (11, 23, 47):
torch.manual_seed(seed)
teacher = Ranker(input_dim, 128).to(DEVICE).eval()
for parameter in teacher.parameters():
parameter.requires_grad_(False)
teachers.append(teacher)
student = Ranker(input_dim, 64).to(DEVICE)
optimizer = torch.optim.AdamW(student.parameters(), lr=2e-3)
alpha = 0.4 # 真实标签损失权重
beta = 0.6 # 蒸馏损失权重
temperature = 3.0
teacher_weights = [0.40, 0.35, 0.25]
for epoch in range(20):
optimizer.zero_grad()
student_logits = student(x)
hard_loss = F.cross_entropy(student_logits, y)
with torch.no_grad():
teacher_probs = []
for teacher in teachers:
logits = teacher(x)
teacher_probs.append(F.softmax(logits / temperature, dim=-1))
merged_teacher_prob = sum(
weight * prob
for weight, prob in zip(teacher_weights, teacher_probs)
)
student_log_prob = F.log_softmax(student_logits / temperature, dim=-1)
soft_loss = F.kl_div(
student_log_prob,
merged_teacher_prob,
reduction="batchmean",
) * (temperature ** 2)
loss = alpha * hard_loss + beta * soft_loss
loss.backward()
optimizer.step()
if (epoch + 1) % 5 == 0:
accuracy = (student_logits.argmax(dim=1) == y).float().mean().item()
print(
f"epoch={epoch + 1:02d} loss={loss.item():.4f} "
f"hard={hard_loss.item():.4f} soft={soft_loss.item():.4f} "
f"accuracy={accuracy:.3f}"
)
print(f"student parameters: {sum(p.numel() for p in student.parameters()):,}")
在真实系统中,可以把示例中的 x 替换成离线生成的查询—职位特征,把三个教师替换成已部署或已训练完成的模型,并把 merged_teacher_prob 保存为蒸馏数据集。这样学生训练阶段不必反复运行昂贵教师,能够降低训练资源和迭代时间。
真正的工程难点不在公式里
多教师蒸馏听起来像是简单的加权平均,但生产系统通常需要处理几个边界问题。
教师输出可能互相冲突
如果不同教师对同一职位给出差异很大的分数,直接平均可能产生一个没有明确语义的目标。可以按流量、查询类型、样本质量或教师在验证集上的表现设置动态权重,也可以先对不同教师的输出做校准。
教师能力不等于线上目标
教师模型可能使用了更丰富、更昂贵、甚至线上无法实时获得的特征。蒸馏时必须明确哪些特征在学生服务阶段可用,否则会出现训练指标很好、上线效果下降的特征泄漏问题。
0.6B 仍然需要严格的服务设计
参数量下降并不自动等于端到端延迟下降。实际收益还取决于:
- 模型是否支持高效批处理;
- 特征读取是否成为主要瓶颈;
- 推理框架是否支持目标硬件;
- 模型输出是否需要额外重排或规则处理;
- 线上流量和候选集规模是否与离线评估一致。
因此,“8 倍更快”应当结合具体训练流程、硬件、数据规模和基线来理解,而不能直接解释成所有线上请求都快 8 倍。
采用前的一份检查清单
如果要把这类方案用于搜索或推荐系统,可以按以下顺序落地:
- 先定义学生模型的边界:明确它负责召回、粗排还是精排,以及线上允许的延迟和成本。
- 拆分教师能力:分别评估语义、行为、结构化匹配等教师信号,而不是默认教师越多越好。
- 建立蒸馏数据集:保存样本特征、真实标签、教师输出和生成时间,便于复现和回滚。
- 比较融合方式:测试固定加权、按场景加权、教师择优和校准后融合。
- 同时看离线与线上指标:关注 NDCG、MRR、点击或申请率、覆盖率、延迟和资源成本。
- 保留真实标签监督:不要让学生完全服从教师,避免放大教师已有的偏差。
- 做漂移监控:职位内容、用户行为和查询分布变化后,旧教师信号可能不再适用。
多教师知识蒸馏的价值,在于把“大模型负责理解、小模型负责规模化执行”变成一条可操作的工程流水线。LinkedIn 选择 0.6B 参数的学生排序模型,也说明模型压缩的目标不是追求最小参数,而是在相关性、训练效率、服务成本和可维护性之间找到一个能上线的平衡点。