☰
用ColBERT做Rerank:从环境搭建到微调评估的实践指南
2026/9/30 3:25:27 网站建设 项目流程

简介:一份面向NLP与深度学习初/中级工程师的rerank模型实践指南,聚焦Sentence Transformers与ColBERT系列模型,系统讲解从环境搭建、bi-encoder与cross-encoder调用,到基于llamaindex接入网易有道embedding/rerank模型,再到微调与MTEB评估的完整链路。资源为PDF格式,共1个文件,压缩包约246KB,内容紧凑且代码示例丰富,适合已有一定基础、希望落地检索排序场景的读者。目前已有118人学习。与泛泛理论不同,这份资料直接给出pip/conda安装命令、模型加载与打分代码、llamaindex后处理串联方式,并展开讲解autotrain微调步骤及c-mteb评估方法,读者可据此快速复现实验,再迁移到问答、知识库检索等真实项目。

1. 搜索排到 60 名的正确结果,才意识到 rerank 不是锦上添花

做自然语言处理检索类项目的人,大概率都经历过这个场景:向量召回模型用 Sentence Transformers 把 query 和文档编码成向量,在 faiss 里一顿操作,Top 50 里正确答案却排到 58 位。直接拿这个结果去做问答、做知识库召回,体验非常糟。这个阶段真正缺的不是更好的召回模型,而是召回和最终答案之间那层精排。rerank 的职责就是拿更强、更精细的模型,把召回回来的几十条重新排一遍。这里有一个常见的路线选择:用 cross-encoder 逐对打分,精度高但慢;用 ColBERT 的 late interaction,在不太牺牲速度的前提下逼近 cross-encoder 的效果。

这篇文章要把这条路从零走通:装环境、跑最小检索流水线,再理解 ColBERT 为什么适合做 rerank,接着做数据微调和评估。适合正在搭 RAG 检索、做知识库问答、或者觉得当前向量召回效果上限太低的人参考。

2. 先把环境搭平:Sentence Transformers 安装与最小检索流水线

2.1 环境和依赖版本怎么选

rerank 模型本质上还是要跑 Transformer,环境里第一优先是 PyTorch 的 CUDA 版本和 GPU 驱动匹配。我一般先建一个干净的 Python 3.10 虚拟环境,再用 pip 装一圈核心依赖,顺序比较重要:

python -m venv venv source venv/bin/activate pip install --upgrade pip pip install torch --index-url https://download.pytorch.org/whl/cu118 pip install sentence-transformers transformers datasets faiss-cpu

提示:如果机器上 GPU 显存 8G 以下,先装 faiss-cpu 起步,评估阶段 CPU 已经完全够用。真正上生产再考虑 faiss-gpu。

这段命令有几个细节值得说。PyTorch 单独先装,是为了避免 pip 自动拉一个 CPU 版本的 torch 进去,一旦 sentence-transformers 的依赖解析把它覆盖成 CPU 版,后面训练速度直接劝退。faiss-cpu 和 faiss-gpu 不能共存,同一个环境里二选一。sentence-transformers 自带model.encode()调用,底层走的是 PyTorch,不需要额外装 sklearn,但它自带的评估器会用到,最好顺手pip install scikit-learn。

装完跑一条命令确认 GPU 可见:

import torch print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0))

如果输出 False,大概率是 torch 版本和驱动不匹配,或者装的 CUDA 版不对。先把 torch 卸了重装对应版本,别急着调模型。

2.2 bi-encoder 做召回的最小脚本

Sentence Transformers 生态里最标准的召回方式,是把 query 和 doc 分别编码成单个向量,然后通过余弦相似度或者 faiss 做最近邻检索。这个方案快,但信息压缩严重,query 里某个关键词的细节很可能在平均池化时被稀释掉。先跑通这条基线,后面和 ColBERT 对比才有参照。

from sentence_transformers import SentenceTransformer, util model = SentenceTransformer("BAAI/bge-large-zh-v1.5") queries = ["营业执照经营范围变更需要什么材料"] docs = [ "企业登记管理办法 第三章 变更登记", "营业执照上的经营范围如何申请变更", "公司章程修正案备案流程", ] # 编码时对 query 和 doc 单独处理,query 指令是可选的 query_emb = model.encode(queries, normalize_embeddings=True) doc_emb = model.encode(docs, normalize_embeddings=True) # 直接使用余弦相似度矩阵 scores = util.cos_sim(query_emb, doc_emb)[0] top_k = scores.topk(k=2) for idx, score in zip(top_k.indices.tolist(), top_k.values.tolist()): print(f"{docs[idx]}: {score:.4f}")

这段代码里,normalize_embeddings=True这个参数决定了后面能不能直接用内积代替余弦距离。faiss 的IndexFlatIP算的是内积,向量归一化之后内积就等于余弦相似度,效率更高。bge 系列模型官方建议在 query 前面加指令后缀,但这里先不加,等微调阶段再看真实数据效果。

跑通之后可以做一个简单的召回验证:把 docs 换成你的知识库切片,query 用真实用户问题。如果 top1 都不是答案,先别急着上 ColBERT,检查切片切得对不对、文档编码是不是把标题和正文截断了。

2.3 cross-encoder 做精排的最小脚本

cross-encoder 的做法是把 query 和 doc 拼成一个长序列输入模型,走完整 Transformer 前向,输出一个相关性分数。这个方案精度最高,但每对都要单独过一遍模型,所以一般只用来精排召回的 50 到 100 条。

from sentence_transformers import CrossEncoder reranker = CrossEncoder("BAAI/bge-reranker-v2-m3") pairs = [[queries[0], doc] for doc in docs] scores = reranker.predict(pairs, batch_size=8) # 把分数降序排,返回原始文档 sorted_idx = scores.argsort(descending=True) for idx in sorted_idx.tolist(): print(f"{docs[idx]}: {scores[idx]:.4f}")

batch_size在这里很关键。predict接收多个 pair 后内部会按 batch 跑,batch 太小 GPU 利用率上不去,太大容易 OOM。一般从 16 起步,看显存逐步往上加。另外注意 rerank 模型和 embedding 模型是两个独立模型,显存占用是两者叠加的,很多 8G 显存的机器在这里开始吃力。

2.4 把召回-精排串起来的参数判断

把上面两个步骤串成完整流水线,中间有个参数需要反复试:召回多少条给精排。

def retrieve_and_rerank(query, top_k_recall=50, top_k_rerank=10): # 召回阶段取 50 条,给精排的候选越多,查全率越高 recall_results = search_index(query_emb, top_k_recall) # 精排阶段只保留 top 10 pairs = [[query, doc] for _, doc in recall_results] scores = reranker.predict(pairs) reranked = [(doc, score) for _, (doc, score) in zip(recall_results, zip(recall_results, scores))] return sorted(reranked, key=lambda x: x[1], reverse=True)[:top_k_rerank]

top_k_recall建议在 50 到 100 之间。如果召回只有 20 条,正确答案压根进不了候选,精排再强也没用。如果召回 200 条,精排耗时线性增长,延迟翻倍但收益很小。判断标准很简单:拿一批真实 query,统计正确答案在召回结果里的平均位置。平均位置在 30 左右,top_k_recall 就设 80 并配合精排,留出缓冲。如果模型运行在 CPU 上,cross-encoder 精排 50 条可能要几百毫秒,这时候可以考虑用后面章节说的 ColBERT,它的索引预计算能省下不少时间。

3. ColBERT 的 late interaction 机制与落地推理

3.1 为什么中间要做 token 级交互

bi-encoder 的一个痛点是:query 和 doc 各有各的向量,两者只做一次向量内积,query 里的“变更”这个词可能被其他词向量淹没。cross-encoder 解决得更彻底,但代价是每对都完整交互一遍。ColBERT 的 late interaction 是两者的折中:它不让 query 和 doc 在输入层就拼一起,而是各自编码,保留 token 级别的向量,最后算相似度时让 query 的每个 token 向量去和 doc 的所有 token 向量逐一算点积,取每个 query token 对应的最大分值再求和。

Sim(q, d) = sum_{i in query tokens} max_{j in doc tokens} E(q_i) · E(d_j)

这个 MaxSim 操作有几个直接的好处。第一,query 里每个词都能在文档里找到和自己最匹配的那一个 token,不会因为平均池化丢了细节。第二,doc 的 token 向量可以提前编码并建索引,不用像 cross-encoder 那样在线逐对计算,运行阶段速度接近 bi-encoder。这也是为什么 ColBERT 被广泛用在 rerank 层:比 bi-encoder 准,比 cross-encoder 快,算是检索精度和成本之间的平衡点。

3.2 用 RAGatouille 把 ColBERT 索引建起来

ColBERT 模型本身有官方仓库,但对工程实践来说,RAGatouille 这个封装可以直接用,它把索引构建、检索、打分封装成一个模型对象。安装方式和模型加载代码如下:

pip install ragatouille
from ragatouille import RAGPretrainedModel colbert = RAGPretrainedModel.from_pretrained("colbert-ir/colbertv2.0")

这里下载的 checkpoint 是 ColBERT v2 在 MS MARCO 上训练的版本。注意 from_pretrained 会从 HuggingFace Hub 拉权重,网络环境不好时容易中断。国内机器可以先用 hf-mirror 之类的镜像把模型下到本地缓存,再改成本地目录路径加载。

索引构建是这个方案里最需要理解的一个环节,看下面的代码:

docs = [ "企业登记管理办法 第三章 变更登记", "营业执照上的经营范围如何申请变更", "公司章程修正案备案流程", "个体工商户登记管理办法 第十条", ] colbert.index( index_name="legal_chat", docs=docs, max_document_length=180, overwrite_index=True, )

index_name会生成一个本地索引目录,这个目录里存的是每个文档 token 向量的分片文件、文档 id 映射和元信息。max_document_length是控制在编码时每个文档最多保留多少 token,ColBERT 编码时会按这个值截断,超过的部分直接丢掉。这个参数直接决定索引体积和后续检索时计算量,常见设置在 150 到 220 之间。如果文档本身很长,比如法律条款原文,建议先用切片逻辑切成 200 字左右的块,再喂给 ColBERT,而不是把演讲级别的长文整个塞进去。

3.3 检索接口与得分可解释性

索引建完之后,检索接口和 faiss 不太一样,它不是返回距离,而是返回一个可解释的分数:

results = colbert.search(query="营业执照经营范围变更需要什么材料", k=3) for hit in results: print(hit["rank"], hit["score"], hit["content"])

这里k是返回的候选条数,注意是“重排后的 top k”,不是召回的候选数。RAGatouille 内部先做向量召回,再对召回结果做 late interaction 重排。如果你只需要精排候选,可以把k设大一点,比如 50,拿到这个分数列表后,再根据业务规则做最终的 top 10 截断。colbert 返回的score是 MaxSim 累加值,不是概率,所以跨 query 的分数不可比,不要拿它和 cross-encoder 的 sigmoid 输出混在一起排序。

3.4 nbits 量化与 doc_maxlen 对索引体积的影响

ColBERT 的索引默认存的是 fp32 的 token 向量,一个文档如果 180 个 token、每个向量 128 维,单条文档就要 90KB 左右。一万条文档就是接近 1G 的索引。这是新手最容易忽视的点。

RAGatouille 的index方法支持nbits参数,比较常用的是 2-bit 量化,索引体积缩小到原来的 1/16 左右,检索质量下降 1% 以内。

colbert.index( index_name="legal_chat_quantized", docs=docs, max_document_length=180, nbits=2, overwrite_index=True, )

nbits的选项一般是 2 和 4。2-bit 体积最小,4-bit 精度更稳。实际项目里建议先用 4-bit 跑通评估集,确认 MRR 指标达标后再把 nbits 降到 2 做压测。如果文档数量超过十万条,这一步能节省几十 G 的磁盘和内存,同时检索延迟也会明显下降。

另一个可以调的是doc_maxlen。业务文档平均长度只有 100 token,却设成 300,等于一半索引都是 padding 出来的空向量。建议先抽样统计文档 token 长度分布,取 85 分位作为max_document_length的初始值。

4. 微调训练:数据组织、损失函数与训练参数

4.1 训练数据长什么样

Sentence Transformers 生态里,微调数据最常见的三种形态是 pair、triplet 和带分数的 pair。

pair 格式最简单,一条数据只包含一个 query 和一个 positive doc,例如:

InputExample(texts=["营业执照经营范围变更需要什么材料", "经营范围变更登记提交材料规范"], label=1)

triplet 格式会多一个 negative doc,让模型学会拉开正样本和负样本的分数差距:

InputExample(texts=[ "营业执照经营范围变更需要什么材料", "经营范围变更登记提交材料规范", "企业年度报告报送方式" ])

带分数的 pair 适合你有业务侧的人工标注,例如运营给每条 query-doc 对打了 0 到 1 的相关性分,这种情况直接用回归或排序损失。

数据怎么挖负样本,是决定微调效果上限的关键。常规做法是先用 BM25 或之前搭好的 bi-encoder 检索,把分数排在第 10 到 50 名、但业务侧确认不相关的文档挖出来作为 hard negative。只用随机负样本会让模型学到“这俩文档主题不同”这种粗粒度能力,对真实场景中“字面很相关但内容不对”的情况毫无帮助。hard negative 控制在正样本数量的 1 到 3 倍之间。

4.2 三种常见损失函数的适用边界

sentence-transformers 官方实现了多套损失,实际微调 rerank 时常用只有三种:MultipleNegativesRankingLoss、CoSENTLoss和SoftmaxLoss。

MultipleNegativesRankingLoss适合只有正样本 pair、没有人工打分的情况,它会自动把 batch 内的其他 query 对应的正样本当作负样本,所以 batch size 对效果影响很大,一般从 32 起步,越大越稳定。

CoSENTLoss需要每个 pair 带一个相似度分数标签,适合有分数标注的数据。它和CosineSimilarityLoss的区别在于用排序不等式约束,收敛更稳。

SoftmaxLoss适合分类式构造,比如把 query 和 4 个文档配对,其中 1 个正确,3 个错误,让模型把这 5 类分出来。

选型逻辑很简单:数据里没有负样本但有大量 query-positive 对,用MultipleNegativesRankingLoss;有分数标签,用CoSENTLoss;有采样好的多分类候选,用SoftmaxLoss。不建议一开始就把三种 loss 叠加当成万能药,先跑通一个,再逐步加。

4.3 一个能跑的训练脚本框架

下面这个是完整的最小微调脚本,基于 sentence-transformers 的fit方法,可以直接改数据跑:

from sentence_transformers import SentenceTransformer, InputExample, losses from torch.utils.data import DataLoader model = SentenceTransformer("BAAI/bge-large-zh-v1.5") train_data = [ InputExample(texts=["营业执照经营范围变更需要什么材料", "经营范围变更登记提交材料规范"]), InputExample(texts=["公司章程修正案要备案吗", "公司章程修正案备案材料清单"]), InputExample(texts=["企业年报逾期怎么办", "企业年度报告公示操作指引"]), ] train_dataloader = DataLoader(train_data, shuffle=True, batch_size=32) loss = losses.MultipleNegativesRankingLoss(model=model) model.fit( train_objectives=[(train_dataloader, loss)], epochs=5, warmup_steps=200, output_path="./fine-tuned-embedding", show_progress_bar=True, evaluation_steps=500, )

这个脚本里有两个参数值得单独说明。warmup_steps一般设置为总训练步数的 10%,作用是让学习率在前 200 步缓慢爬升,防止开头把预训练权重冲坏。evaluation_steps是每隔多少步跑一次评估,如果设了 500,训练数据少于 500 条会直接跳过评估,这个参数不要大于训练步数的一半。

4.4 微调产出物怎么替换到流水线

model.fit的output_path目录下会保存模型权重和配置,替换的时候只需要把这行代码换掉:

model = SentenceTransformer("./fine-tuned-embedding")

替换之后要做一次与微调前完全相同的数据集评估,对比 MRR 或 Recall@k。如果微调后指标变差,不要急着调损失函数,先看是不是负样本来源和评估集分布差异太大。另外,微调阶段建议只更新最后一两层,也就是在fit中传入model前,先冻结前面的层:

for param in model.parameters(): param.requires_grad = False for param in model[1].parameters(): param.requires_grad = True

sentence-transformers的模型结构里model[1]是 pooling 层,model[0].auto_model才是 Transformer 本体。如果你对 PyTorch 不熟,最简单做法是不冻结,直接全量微调,但数据量少于几千对时效果反而不稳定。

5. 避坑:微调与评估阶段最常踩的五个坑

5.1 训练 loss 震荡不降,负样本难易没控制好

现象:loss 前期下降很快,训练到一半开始剧烈震荡,甚至逐步回升。原因:负样本挖得太简单或者太难。具体来说,如果 batch 内随机采样到的 negative 都是完全不同主题的文档,模型很快就学会粗略区分,loss 降到一定程度就失去梯度。如果 hard negative 全部是和 query 极相似的擦边文档,模型又会被带偏。解决思路是混合负样本:2/3 用中等难度的 BM25 负样本,1/3 用随机采样,这样 loss 会比较平稳。另外把MultipleNegativesRankingLoss默认的 temperature 从 0.05 调整到 0.02,也能缓解训练后期分数被压得过平的问题。

5.2 ColBERT 索引磁盘体积爆炸,比原文档大几十倍

现象:一万条文档建出来的索引目录动辄几个 G 到十几个 G。原因:nbits没设置,默认用 fp32 存储所有 token 向量,并且max_document_length设置过高导致大量 padding 向量也被存了。解决:给index()方法传nbits=2,同时把max_document_length按语料 85 分位调整。做完这两步,索引体积通常会缩小到原来的 1/8 到 1/20,MRR 损失在 1% 上下。这也是量化在检索链路里最划算的一笔投入。

5.3 稍微调大 batch_size 就 OOM

现象:batch_size=32没问题,改成 64 直接 CUDA out of memory。原因:不是单纯整数翻倍,而是 Padding 导致的计算量膨胀。文本长度不一样的时候,一个 batch 里所有样本都会 pad 到最长那条的长度,如果最长的那条是 512 token,其他 30 条都是 64 token,等于白算了 6 倍。解决:用动态 batch 策略,按文本长度分组,长度相近的放同一个 batch。同时启用model.fit(use_amp=True)来跑混合精度,显存占用能降 40%。做了这两步还 OOM,再把max_seq_length从 512 降到 256。

5.4 GPU 利用率上不去,训练速度像在跑 CPU

现象:nvidia-smi看显存占满了,但util只有 20%,一个 epoch 要跑半天。原因:数据加载是瓶颈,尤其当数据是中文长文本时,tokenizer的预处理全部卡在 CPU 端,GPU 一直在等数据。解决:DataLoader里设置num_workers=4或更高,并确认fit传入的数据对象不是生成器而是完整 list,否则 worker 没法预取。另外把show_progress_bar开起来,通过 it/s 判断速度变化,这个指标比肉眼盯着 GPU 利用率更直接。

5.5 微调后离线评估反而变差,评估集疑似被污染

现象:模型在训练数据上效果很好,但在独立的业务 query 上 MRR 下降。原因:微调数据的 query 分布和线上差异太大,比如训练数据来自搜索词,评估数据是口语化长句,模型被微调“带偏”了。另外如果评估集是从训练集里抽出来的,或者负样本挖掘时把评估集的文档也送进了训练负样本里,评估分数就会失真。解决:评估集必须完全独立,且在挖负样本时显式排除评估集文档。给评估集加一层 filter:

eval_doc_ids = set([d["id"] for d in eval_set]) def is_valid_negative(doc_id): return doc_id not in eval_doc_ids

这个看起来无关紧要的细节,往往是模型上线后效果崩掉的真正原因,值得在做评估前先确认一遍。

6. 评估结果怎么验证:两个指标和一个可抄的脚本

rerank 模型够不够好,不能只靠肉眼抽查几条排序结果。离线评估至少要覆盖两个指标:MRR(Mean Reciprocal Rank)和 Recall@k。MRR 关注正确结果的排位有多靠前,Recall@k 关注正确结果有没有进入前 k 条。前者更贴合用户只看前几名的场景,后者更适合对查全率有硬性要求的管道。

def evaluate_mrr(qrels, results, k=10): scores = [] for query_id, gold_doc_id in qrels.items(): ranked = [doc_id for doc_id, _ in results[query_id][:k]] if gold_doc_id in ranked: scores.append(1.0 / (ranked.index(gold_doc_id) + 1)) else: scores.append(0.0) return sum(scores) / len(scores) def evaluate_recall_at_k(qrels, results, k=10): hits = 0 for query_id, gold_doc_id in qrels.items(): ranked = [doc_id for doc_id, _ in results[query_id][:k]] if gold_doc_id in ranked: hits += 1 return hits / len(qrels)

qrels是 query 到标准答案文档的映射,results是模型对每个 query 输出的 top-k 结果列表。评估的时候把 k 分别设成 1、5、10 观察曲线,如果 Recall@10 高但 MRR 低,说明答案经常出现在第 6 到 10 名附近,这种情况优先调精排模型或者加大召回候选数,而不是反复调向量模型。如果 MRR@1 高但 Recall@10 低,说明模型前几名很准但覆盖不够,需要回到召回阶段补切片或加索引。这个脚本跑完之后,记得把评估结果和微调前对照打印出来,最好顺带记录每条 query 的响应时间,评估不只是看效果,还要看 rerank 层有没有把整个检索链路拖慢。我之前有次微调后 MRR 涨了 3 个点,但 p95 延迟从 80ms 涨到 300ms,最后还是回退了版本。检索系统是延迟敏感型,评估指标里一定要预留响应时间这一栏。希望这些记录能帮你少踩几个我用一个又一个不眠夜换来的坑。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询