LinkedIn 如何用多教师知识蒸馏把 AI 求职排序训练提速 8 倍

2026-09-11 24 预计阅读时间: 1 分钟
来源: infoq.com AI 摘要 Original link

Disclaimer: This article is an AI-assisted summary. Read it together with the original source when precision matters. The summary may omit context, version differences, or edge cases and is not official documentation.

预计阅读时间:11 分钟

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 损失乘以 ,避免温度变化导致梯度规模失控。

一个可以运行和改造的 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 倍。

采用前的一份检查清单

如果要把这类方案用于搜索或推荐系统,可以按以下顺序落地:

  1. 先定义学生模型的边界:明确它负责召回、粗排还是精排,以及线上允许的延迟和成本。
  2. 拆分教师能力:分别评估语义、行为、结构化匹配等教师信号,而不是默认教师越多越好。
  3. 建立蒸馏数据集:保存样本特征、真实标签、教师输出和生成时间,便于复现和回滚。
  4. 比较融合方式:测试固定加权、按场景加权、教师择优和校准后融合。
  5. 同时看离线与线上指标:关注 NDCG、MRR、点击或申请率、覆盖率、延迟和资源成本。
  6. 保留真实标签监督:不要让学生完全服从教师,避免放大教师已有的偏差。
  7. 做漂移监控:职位内容、用户行为和查询分布变化后,旧教师信号可能不再适用。

多教师知识蒸馏的价值,在于把“大模型负责理解、小模型负责规模化执行”变成一条可操作的工程流水线。LinkedIn 选择 0.6B 参数的学生排序模型,也说明模型压缩的目标不是追求最小参数,而是在相关性、训练效率、服务成本和可维护性之间找到一个能上线的平衡点。


相关推荐