判别式语言模型在检索系统中的应用:从双塔到直接打分
在实际的搜索和推荐系统中,如何高效、准确地从海量候选集中召回相关项目(Item)一直是个核心挑战。传统的双塔模型通过将查询(Query)和项目(Item)分别映射到同一向量空间进行相似度计算,虽然高效,但往往需要为每个项目生成一个唯一的标识符(Item ID)并预计算其向量。当项目库动态变化或项目本身是复杂文本时,这种基于ID的预计算方式在灵活性上会受到限制。近期,Meta的一项研究提出了一种新思路:直接使用判别式语言模型(Discriminative Language Model)作为检索器,它能够直接对查询和候选项目文本进行打分,从而绕过了生成Item ID和预计算向量的步骤。
这种方法的核心在于,模型不再学习将项目压缩成一个静态ID或向量,而是作为一个“判别器”,直接评估给定查询下,某个项目文本作为正确答案的可能性。这听起来像是把检索任务重新定义为了一个文本匹配或分类问题。对于从事搜索、推荐、问答系统开发的工程师和研究者来说,理解这种范式转变背后的动机、实现路径以及其与经典双塔模型的优劣对比,对于技术选型和架构演进至关重要。本文将深入解析这一技术路径,从概念原理、模型架构、训练方法到实践中的关键考量,提供一个可理解、可评估的技术视角。
1. 理解判别式语言模型作为检索器的核心思想
要理解这项技术,首先需要厘清几个关键概念:生成式模型、判别式模型、双塔检索以及它们在此项工作中的应用方式。
1.1 生成式与判别式模型的区别
在机器学习中,生成式模型(如GPT系列)学习的是数据的联合概率分布 P(X, Y),其目标是建模数据是如何“生成”的,因此可以用于生成新的数据样本。而判别式模型(如BERT、分类器)学习的是条件概率分布 P(Y|X),其目标是直接学习在给定输入X的情况下,输出Y的边界或概率,更专注于“判别”或“分类”任务。
传统的基于BERT的双塔检索模型,本质上也是一种判别式模型的应用:它通过对比学习等方式,训练模型将查询和正样本项目的向量拉近,与负样本的向量推远。然而,它的输出是一个“向量”,检索时需要通过向量相似度计算(如点积、余弦相似度)来完成。而Meta论文中提出的方法,是将判别式语言模型的输出直接用于“打分”。
1.2 从“向量检索”到“直接打分”的范式转变
在双塔架构中,流程通常是:
- 离线:为所有项目(Item)生成ID,并通过项目塔(Item Tower)模型计算其向量表示,存入向量数据库。
- 在线:收到用户查询(Query)后,通过查询塔(Query Tower)模型计算查询向量。
- 检索:在向量数据库中执行近似最近邻搜索(ANN),找出与查询向量最相似的项目向量,返回对应的Item ID。
这个过程强依赖于预计算的Item向量。而判别式语言模型作为检索器的思路则截然不同:
- 模型角色转变:模型本身就是一个打分函数
f(query, item_text)。 - 输入输出:输入是原始的查询文本和候选项目的原始文本(或结构化文本表示),输出是一个标量分数,直接表示该item与query的相关性。
- 检索过程:对于每个查询,需要将它与所有候选项目的文本(或一个经过筛选的子集)逐一输入模型进行打分,然后按分数排序。这听起来计算量巨大,但可以通过高效的模型设计、负采样策略和推理优化来缓解。
这种方法的优势在于灵活性:项目库可以动态增删,无需重新训练模型来生成新的Item ID或向量,只需将新项目的文本加入候选池即可。同时,它能够充分利用项目的完整文本信息,而不是被压缩到一个固定维度的向量中。
1.3 与生成式检索(GENRE)的对比
另一种绕过传统检索范式的方法是生成式检索,例如GENRE(Generative ENtity REtrieval)模型。它直接将检索任务视为一个序列生成问题,模型被训练来直接生成目标实体(Item)的标识符(如标题、ID)。虽然也避免了显式的向量相似度计算,但它属于生成式范式。
本文讨论的判别式方法与之关键区别在于:
- 目标不同:生成式模型学习
P(item_id | query),判别式模型学习P(relevance_score | query, item_text)。 - 输出不同:生成式输出是文本(ID),需要处理生成重复、未知标识符等问题;判别式输出是分数,更直接,且天然支持对已知候选集进行排序。
- 灵活性:判别式方法可以轻松处理项目文本描述的变化,而生成式方法如果项目文本发生变化,其对应的生成目标可能需要调整。
2. 模型架构与训练方法设计
要将一个判别式语言模型(如BERT、RoBERTa)改造成高效的检索器,需要在模型架构、输入处理和训练目标上进行特殊设计。
2.1 模型架构:编码器与打分头
通常采用一个预训练的语言模型编码器(如BERT)作为主干网络。其关键设计在于如何将查询和项目文本组合,并产生一个相关性分数。
1. 输入表示:查询文本和项目文本不会被分别编码成两个向量,而是被拼接成一个序列,作为编码器的联合输入。格式通常如下:
[CLS] Query Text [SEP] Item Text [SEP]这种格式让模型能够充分捕捉查询和项目之间的交叉注意力(Cross-Attention),这是双塔模型不具备的能力,双塔模型在编码阶段查询和项目是相互独立的。
2. 打分头(Scoring Head):编码器输出[CLS]位置的隐藏状态(或整个序列的池化结果)被输入到一个简单的打分头,通常是一个线性层(Linear Layer),将高维向量映射为一个标量分数。
import torch import torch.nn as nn from transformers import AutoModel, AutoTokenizer class DiscriminativeRetriever(nn.Module): def __init__(self, model_name='bert-base-uncased'): super().__init__() self.encoder = AutoModel.from_pretrained(model_name) self.scorer = nn.Linear(self.encoder.config.hidden_size, 1) # 打分头 def forward(self, query_input_ids, query_attention_mask, item_input_ids, item_attention_mask): # 拼接查询和项目输入 input_ids = torch.cat([query_input_ids, item_input_ids], dim=1) attention_mask = torch.cat([query_attention_mask, item_attention_mask], dim=1) # 通过编码器 outputs = self.encoder(input_ids=input_ids, attention_mask=attention_mask) # 取[CLS]位置的输出 cls_output = outputs.last_hidden_state[:, 0, :] # 计算分数 score = self.scorer(cls_output).squeeze(-1) # 形状: (batch_size,) return score # 示例:初始化模型和分词器 model = DiscriminativeRetriever() tokenizer = AutoTokenizer.from_pretrained('bert-base-uncased')2.2 训练目标:对比学习与列表级损失
训练的核心目标是让模型学会给相关(正样本)的(query, item)对打高分,给不相关(负样本)的打低分。常用的损失函数包括:
1. 二元交叉熵损失(Binary Cross-Entropy Loss):将任务视为一个二分类问题(相关/不相关)。对于每个(query, positive_item)对,构造若干个(query, negative_item)对,然后使用sigmoid函数将模型输出分数转换为概率,计算交叉熵损失。
# 假设 scores 是模型对一批 (query, item) 对输出的分数 # labels 是二分类标签,1表示正样本,0表示负样本 loss_fn = nn.BCEWithLogitsLoss() # 内部包含sigmoid loss = loss_fn(scores, labels.float())这种方法的挑战在于负样本的构造。高质量的负样本(困难负样本)对模型性能至关重要。
2. 对比损失(Contrastive Loss)或 InfoNCE Loss:更常见于检索任务。对于一个查询q,有一个正样本i+和多个负样本{i1-, i2-, ..., iN-}。模型对所有(q, i)对进行打分,然后计算softmax交叉熵损失,目标是让正样本的分数远高于负样本。
# scores: 形状为 (batch_size, num_candidates),其中每一行是一个query对应其正样本和多个负样本的分数 # 假设每行的第一个分数是正样本的分数 positive_scores = scores[:, 0].unsqueeze(1) # (batch_size, 1) # 计算softmax概率(温度参数tau用于平滑) logits = scores / tau # 标签:每行的第一个位置是正类 labels = torch.zeros(scores.size(0), dtype=torch.long).to(scores.device) loss = nn.CrossEntropyLoss()(logits, labels)这种列表级的损失函数迫使模型在候选集中进行区分,更符合实际检索的排序场景。
2.3 知识蒸馏的应用
论文中提到“知识蒸馏”,这可能是提升模型性能的关键技术。具体来说,可以用一个更大、更复杂的模型(教师模型)来为(query, item)对生成软标签(soft scores),然后让当前要训练的轻量级模型(学生模型)去拟合这些软标签。
为什么需要知识蒸馏?
- 数据效率:教师模型可能从海量无监督或弱监督数据中学习到了更丰富的语义匹配知识。
- 标签平滑:软标签提供了比硬标签(0/1)更丰富的监督信号,例如,一个项目可能与查询“部分相关”,得分为0.7。
- 模型压缩:最终部署的判别式检索器需要极高的推理速度,因此通常是一个较小的模型。通过知识蒸馏,小模型可以继承大模型的能力。
蒸馏损失通常结合了硬标签损失和软标签损失:
# student_scores, teacher_scores 分别是学生和教师模型对同一批输入的打分 hard_loss = contrastive_loss(student_scores, hard_labels) # 使用真实标签的对比损失 # 使用KL散度衡量学生输出分布与教师输出分布的差异 soft_loss = nn.KLDivLoss(reduction='batchmean')( F.log_softmax(student_scores / T, dim=1), F.softmax(teacher_scores / T, dim=1) ) total_loss = alpha * hard_loss + (1 - alpha) * soft_loss * (T**2) # T是温度,alpha是权重3. 实践中的关键考量与实现步骤
将判别式语言模型应用于实际检索场景,会面临效率、负采样、部署等一系列工程挑战。
3.1 效率挑战与优化策略
最直接的挑战是:对于每个查询,如何避免与百万甚至千万级别的候选项目逐一计算分数?
1. 召回-精排两阶段架构:这是工业界标准做法,判别式模型通常用于“精排”阶段。
- 召回阶段:使用传统的双塔模型、倒排索引或轻量级ANN方法,快速从全量库中筛选出Top K(例如1000个)候选。
- 精排阶段:将查询与这K个候选项目的文本,输入判别式模型进行精细打分和重排序。 这样,判别式模型只需要处理K个候选,而不是全量库。
2. 模型与推理优化:
- 模型轻量化:使用知识蒸馏训练更小的模型(如TinyBERT、DistilBERT),或使用模型剪枝、量化技术。
- 批处理与硬件加速:在GPU上对
(query, K个item)进行批量并行打分。由于输入是[CLS] Q [SEP] I [SEP],可以构建一个批量为[Q+I1, Q+I2, ..., Q+Ik]的输入。 - 缓存与索引:虽然项目文本会变,但查询侧的部分计算或项目的某些固定特征可以尝试缓存。
3.2 负样本采样策略
训练数据的质量,尤其是负样本的质量,直接决定模型区分好坏的能力。
常见负样本来源:
- 随机负样本:从全体项目中随机抽取。简单但质量低,模型容易学习。
- 批量内负样本:在一个训练批次中,将其他正样本对应的项目作为当前查询的负样本。这是对比学习中的常用技巧。
- 困难负样本:使用上一代检索模型或双塔模型,为每个查询召回一批得分较高但不是正样本的项目。这些是模型容易混淆的样本,对提升模型性能至关重要。
- 人工构造负样本:通过规则或启发式方法构造与查询相似但无关的项目文本。
一个鲁棒的训练流程通常会混合使用多种负样本。
3.3 端到端实现步骤示例
假设我们有一个(query, positive_item_title)的配对数据集,目标是训练一个用于文章标题检索的判别式模型。
步骤1:环境准备与数据预处理
# 环境依赖 # transformers, torch, datasets, tqdm, numpy, pandas import pandas as pd from datasets import Dataset # 假设数据格式:csv文件,包含 query, pos_title 两列 df = pd.read_csv('retrieval_data.csv') # 构建训练样本:为每个query构造负样本(这里简单使用批量内负样本,实际需更复杂策略) dataset = Dataset.from_pandas(df[['query', 'pos_title']])步骤2:定义数据加载与负采样
from torch.utils.data import DataLoader import random def collate_fn(batch, tokenizer, max_length=128): queries = [item['query'] for item in batch] pos_titles = [item['pos_title'] for item in batch] # 简单的批量内负采样:将同一batch内其他样本的正标题作为负样本 neg_titles = [] for i in range(len(batch)): # 排除自身 candidates = pos_titles[:i] + pos_titles[i+1:] neg_titles.append(random.choice(candidates) if candidates else pos_titles[i]) # 防错 # Tokenize 所有文本对 pos_pairs = [f"{q} [SEP] {t}" for q, t in zip(queries, pos_titles)] neg_pairs = [f"{q} [SEP] {t}" for q, t in zip(queries, neg_titles)] # 编码 pos_encodings = tokenizer(pos_pairs, truncation=True, padding='max_length', max_length=max_length, return_tensors='pt') neg_encodings = tokenizer(neg_pairs, truncation=True, padding='max_length', max_length=max_length, return_tensors='pt') # 注意:这里简化了,实际训练时一个query会对应多个负样本 return { 'pos_input_ids': pos_encodings['input_ids'], 'pos_attention_mask': pos_encodings['attention_mask'], 'neg_input_ids': neg_encodings['input_ids'], 'neg_attention_mask': neg_encodings['attention_mask'] }步骤3:训练循环核心代码
import torch.optim as optim from tqdm import tqdm device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = DiscriminativeRetriever().to(device) optimizer = optim.AdamW(model.parameters(), lr=2e-5) for epoch in range(num_epochs): model.train() total_loss = 0 for batch in tqdm(train_dataloader): optimizer.zero_grad() # 正样本分数 pos_scores = model( batch['pos_input_ids'].to(device), batch['pos_attention_mask'].to(device), # 注意:这里模型定义需要调整以接受拼接好的输入,上述collate_fn也需要调整。 # 更合理的做法是collate_fn直接输出拼接好的正负样本对。 ) # 负样本分数 neg_scores = model( batch['neg_input_ids'].to(device), batch['neg_attention_mask'].to(device), ) # 计算对比损失 (示例,假设每个query只有一个正样本和一个负样本) # scores: 将正负分数拼接,形状为 (batch_size, 2) scores = torch.stack([pos_scores, neg_scores], dim=1) # 标签:正样本在位置0 labels = torch.zeros(scores.size(0), dtype=torch.long).to(device) loss = nn.CrossEntropyLoss()(scores, labels) loss.backward() optimizer.step() total_loss += loss.item() print(f"Epoch {epoch}, Loss: {total_loss / len(train_dataloader)}")注意:以上代码是高度简化的示意,旨在说明流程。实际实现中,数据批处理、负采样策略、损失函数都需要根据论文和具体任务进行精心设计。
4. 与双塔模型的对比分析与选型建议
判别式语言模型检索器并非要完全取代双塔模型,而是提供了另一种技术选项。理解它们的差异是做出正确技术选型的基础。
4.1 核心差异对比
| 特性维度 | 双塔模型 (Bi-Encoder) | 判别式语言模型 (Cross-Encoder) |
|---|---|---|
| 交互时机 | 编码时无交互,后期交互(通过向量点积)。 | 编码时深度交互(通过Transformer注意力机制)。 |
| Item表示 | 预计算的静态向量。 | 原始的动态文本,每次参与计算。 |
| 检索效率 | 极高。在线阶段只需计算一次查询向量,然后进行快速的ANN搜索。 | 较低。需要与每个候选Item进行联合编码计算,复杂度随候选集线性增长。 |
| 精度潜力 | 相对较低。因为查询和项目在编码阶段是独立的。 | 相对更高。能够进行细粒度的语义匹配和消歧。 |
| 灵活性 | 较低。项目库更新需要重新计算所有向量。 | 极高。项目文本变更无需重新训练模型,直接参与计算即可。 |
| 适用场景 | 海量候选集(>百万)的召回阶段、需要极低延迟的在线服务。 | 中小规模候选集(<10万)的精排阶段、对精度要求极高的场景、项目文本频繁变化的场景。 |
4.2 常见问题与排查路径
在实际应用判别式检索器时,可能会遇到以下典型问题:
问题1:模型训练收敛慢或效果不佳。
- 可能原因:负样本质量太差(全是简单负样本),模型学不到有效的区分能力。
- 排查与解决:
- 检查负样本:随机抽取一些训练样本,人工检查
(query, negative_item)对是否真的不相关。如果很多是弱相关的,模型会困惑。 - 引入困难负样本:使用一个基线模型(如BM25、双塔模型)为每个查询召回一批得分较高的非正样本,加入训练。
- 调整损失函数:尝试不同的温度系数
tau,或结合二元交叉熵损失。 - 验证数据划分:确保训练集和验证集没有信息泄露(例如,同一个项目出现在训练集的正样本和验证集的负样本中)。
- 检查负样本:随机抽取一些训练样本,人工检查
问题2:线上推理延迟过高,无法满足服务要求。
- 可能原因:候选集K太大,或模型本身过于复杂。
- 排查与解决:
- 性能剖析:使用性能分析工具(如PyTorch Profiler)定位耗时瓶颈是在模型前向传播还是数据加载。
- 优化召回阶段:收紧召回阶段的条件,减少进入精排的候选数量K。确保召回模型的质量,避免漏掉好的候选。
- 模型压缩:应用知识蒸馏、剪枝、量化(如INT8量化)技术,缩小模型体积,提升推理速度。
- 硬件与批处理:使用GPU并优化批处理大小,充分利用硬件并行能力。考虑使用TensorRT或ONNX Runtime进行推理优化。
问题3:项目文本过长,导致输入超出模型最大长度。
- 可能原因:BERT类模型通常有512或1024的长度限制。
- 排查与解决:
- 文本截断:优先保留项目标题、关键属性、摘要等核心信息,截断长描述。
- 特征工程:将长文本的关键信息(如实体、主题词)提取出来,拼接成短文本作为模型输入。
- 使用长文本模型:考虑使用支持更长序列的模型,如Longformer、BigBird,但需注意其计算开销。
4.3 最佳实践与扩展方向
最佳实践:
- 两阶段架构:始终坚持“召回+精排”的架构。用双塔、ANN做高效召回,用判别式模型做精准重排序。这是平衡效果和效率的黄金法则。
- 渐进式迭代:不要一开始就用复杂的判别式模型。先从简单的基线(如BM25、双塔)开始,建立评估体系,再逐步引入更复杂的模型进行A/B测试。
- 重视负样本:将至少30%的精力花在构建高质量的负样本库上,包括困难负样本和人工审核的负样本。
- 离线评估先行:在上线前,使用离线评估指标(如Recall@K, NDCG, MRR)充分验证模型效果,并与基线模型对比。
- 监控线上指标:上线后,密切监控点击率(CTR)、转化率等业务指标,以及模型服务的延迟、成功率等技术指标。
扩展方向:
- 多模态检索:判别式框架可以自然扩展。输入不仅是文本,可以拼接图像特征向量、结构化属性特征等,让模型学习跨模态的匹配。
- 端到端学习:将召回和精排模型进行联合训练或深度优化,例如让精排模型为召回模型提供反馈信号。
- 与生成式结合:在问答、对话系统中,可以先使用判别式检索器从知识库中找出最相关的文档片段,再交给生成式模型(如GPT)生成最终答案,构建RAG(Retrieval-Augmented Generation)系统。
- 蒸馏到双塔:利用训练好的高性能判别式模型(教师)去蒸馏一个双塔模型(学生),让学生模型在保持高效检索的同时,逼近教师的精度。
判别式语言模型作为检索器,代表了一种更灵活、更注重深度语义匹配的技术路线。它虽然牺牲了部分效率,但在对精度和灵活性要求高的精排场景、动态项目库场景下展现出独特优势。在实际工程中,将其与成熟的向量检索技术结合,构建分层的检索系统,是当前最务实和有效的方案。理解其原理和实现细节,能帮助我们在面对复杂检索需求时,拥有更多样化的技术武器。
