判别式语言模型在检索系统中的应用:从双塔到直接打分

📅 2026/8/13 6:42:00
判别式语言模型在检索系统中的应用:从双塔到直接打分
在实际的搜索和推荐系统中如何高效、准确地从海量候选集中召回相关项目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的对比另一种绕过传统检索范式的方法是生成式检索例如GENREGenerative 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_namebert-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], dim1) attention_mask torch.cat([query_attention_mask, item_attention_mask], dim1) # 通过编码器 outputs self.encoder(input_idsinput_ids, attention_maskattention_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), dtypetorch.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(reductionbatchmean)( F.log_softmax(student_scores / T, dim1), F.softmax(teacher_scores / T, dim1) ) 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]可以构建一个批量为[QI1, QI2, ..., QIk]的输入。缓存与索引虽然项目文本会变但查询侧的部分计算或项目的某些固定特征可以尝试缓存。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_length128): 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[i1:] 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, truncationTrue, paddingmax_length, max_lengthmax_length, return_tensorspt) neg_encodings tokenizer(neg_pairs, truncationTrue, paddingmax_length, max_lengthmax_length, return_tensorspt) # 注意这里简化了实际训练时一个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(), lr2e-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], dim1) # 标签正样本在位置0 labels torch.zeros(scores.size(0), dtypetorch.long).to(device) loss nn.CrossEntropyLoss()(scores, labels) loss.backward() optimizer.step() total_loss loss.item() print(fEpoch {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%的精力花在构建高质量的负样本库上包括困难负样本和人工审核的负样本。离线评估先行在上线前使用离线评估指标如RecallK, NDCG, MRR充分验证模型效果并与基线模型对比。监控线上指标上线后密切监控点击率CTR、转化率等业务指标以及模型服务的延迟、成功率等技术指标。扩展方向多模态检索判别式框架可以自然扩展。输入不仅是文本可以拼接图像特征向量、结构化属性特征等让模型学习跨模态的匹配。端到端学习将召回和精排模型进行联合训练或深度优化例如让精排模型为召回模型提供反馈信号。与生成式结合在问答、对话系统中可以先使用判别式检索器从知识库中找出最相关的文档片段再交给生成式模型如GPT生成最终答案构建RAGRetrieval-Augmented Generation系统。蒸馏到双塔利用训练好的高性能判别式模型教师去蒸馏一个双塔模型学生让学生模型在保持高效检索的同时逼近教师的精度。判别式语言模型作为检索器代表了一种更灵活、更注重深度语义匹配的技术路线。它虽然牺牲了部分效率但在对精度和灵活性要求高的精排场景、动态项目库场景下展现出独特优势。在实际工程中将其与成熟的向量检索技术结合构建分层的检索系统是当前最务实和有效的方案。理解其原理和实现细节能帮助我们在面对复杂检索需求时拥有更多样化的技术武器。