Chapter 40
5. 交叉编码器的重排序 Cross encoder re ranking
NotebookPython 3 (ipykernel)28 cells
第五章 交叉编码器重排序
本节课,将使用交叉编码器重排序的技术,对检索到的结果进行相关性分析。重排序是一种根据结果与特定查询的相关性来排序和评分的方法。
一、底层原理
1.1 重排序
在得到特定查询的检索结果后,需要将该结果同查询一起输入到一个重排序模型中,使得最相关的结果具有最高的排名。 也就是说,重排序模型会根据查询来为每个结果打分,排名最高的结果就是与该查询最相关的结果。
1.2 交叉编码器
BERT交叉编码器是句子转换器的一种模型,可以同时获取查询和文档,通过一个分类器对传入查询和每个检索到的文档打分,最后输出分数。
二、实现过程
2.1 导入辅助函数
In [1]python · cell 11
python
from helper_utils import load_chroma, word_wrap, project_embeddings
import chromadb.utils.embedding_functions as embedding_functions
from chromadb.utils.embedding_functions import SentenceTransformerEmbeddingFunction
import numpy as np
import warnings
# 忽略 FutureWarning 类型的警告
warnings.filterwarnings("ignore", category=FutureWarning)In [2]python · cell 12
python
# 使用代理可能出现网络问题,将以下端口号1080全部替换成自己的vpn的端口号
import os
# os.environ['HTTPS_PROXY']='http://127.0.0.1:1080'
# os.environ["HTTP_PROXY"]='http://127.0.0.1:1080'
import openai
from openai import OpenAI
from dotenv import load_dotenv, find_dotenv
loaded = load_dotenv(find_dotenv(), override=True)
# 从环境变量中获取 OpenAI API Key 或者直接赋值
API_KEY = os.getenv("API_KEY")
# 如果您使用的是官方 API,就直接用 https://api.siliconflow.cn/v1 就行。
BASE_URL = "https://api.siliconflow.cn/v1"In [3]python · cell 13
python
# chromadb支持的嵌入函数有许多种,这里介绍常用的几种:
# 参考资料:https://docs.trychroma.com/embeddings
# 方式1:默认嵌入函数,需要下载模型,本地计算。英文文本表现不错,中文文本表现一般
# embedding_function = SentenceTransformerEmbeddingFunction()
# 方式2:OpenAI的嵌入函数,直接调用OpenAI的接口,无需下载模型(推荐)
embedding_function = embedding_functions.OpenAIEmbeddingFunction(
api_key=API_KEY,
api_base=BASE_URL,
model_name="BAAI/bge-m3",
dimensions=1024
)
# 方式3:HuggingFace的嵌入函数,需要下载模型,本地计算,对网络要求高
# embedding_function = embedding_functions.HuggingFaceEmbeddingFunction(
# api_key="hf_", # 填入你的 huggingface Access Token
# model_name="jinaai/jina-embeddings-v2-base-zh" # 指定模型
# )
## 中文备选模型
# jinaai/jina-embeddings-v2-base-zh
# GanymedeNil/text2vec-large-chinese
# BAAI/bge-large-zh-v1.5
# BAAI/bge-small-zh-v1.5In [4]python · cell 14
python
chroma_collection = load_chroma(filename='./data/2024年北京市政府工作报告.pdf',
collection_name='beijing_annual_report_2024',
embedding_function=embedding_function,langcode='zh')
chroma_collection.count()Output
1028
2.2 长尾部分的重排序
In [5]python · cell 16
python
# 之前一般设定返回5个结果,现在要求返回10个结果,加入了部分可能有用的的长尾结果
query = "地区生产总值是多少?"
results = chroma_collection.query(query_texts=query, n_results=10, include=['documents', 'embeddings'])
retrieved_documents = results['documents'][0]
for document in results['documents'][0]:
print(word_wrap(document))
print('')Output
全市地区生产总 值增长 5.2%、约 4.4 万亿元 数字经济增加值占地区生产总值比重达 42.9% 人均地区生产总值、全 员劳动生产率、万元地区生产总值能耗水耗等多项指标保持全国省级地区最优水平 提出今年全市经济社 会发展的主要预期目标是:地区生产总值增长 5%左右 居民收入增长与经济增长同步 居民收入增长与经济增长同步 推动经济实现质的有效提升和量的合理增长 一般公共预算收入增长 8.2%、突破 6000 亿元 居民消费价格涨幅 3%左右 一般公共预算收入增长 5%
In [6]python · cell 17
python
# BERT交叉编码器同时渠道查询和文档,通过一个分类器传递,获得一个得分
# 利用该得分作为检索结果的相关性或排名的得分
from sentence_transformers import CrossEncoder
cross_encoder = CrossEncoder('cross-encoder/ms-marco-MiniLM-L-6-v2')In [7]python · cell 18
python
pairs = [[query, doc] for doc in retrieved_documents]
scores = cross_encoder.predict(pairs)
print("分数和相应的文本:")
# 打印每个文档的分数和内容
for score, (query, document) in zip(scores, pairs):
print(f"分数: {score}")
print("文本:")
print(word_wrap(document))
print('') # 在文档间添加空行以便区分Output
分数和相应的文本: 分数: 7.281141757965088 文本: 全市地区生产总 值增长 5.2%、约 4.4 万亿元 分数: 7.037868499755859 文本: 数字经济增加值占地区生产总值比重达 42.9% 分数: 8.211084365844727 文本: 人均地区生产总值、全 员劳动生产率、万元地区生产总值能耗水耗等多项指标保持全国省级地区最优水平 分数: 7.3925018310546875 文本: 提出今年全市经济社 会发展的主要预期目标是:地区生产总值增长 5%左右 分数: 0.8092600107192993 文本: 居民收入增长与经济增长同步 分数: 0.8092600107192993 文本: 居民收入增长与经济增长同步 分数: 3.665830612182617 文本: 推动经济实现质的有效提升和量的合理增长 分数: 3.7360196113586426 文本: 一般公共预算收入增长 8.2%、突破 6000 亿元 分数: -2.139439105987549 文本: 居民消费价格涨幅 3%左右 分数: 1.5210455656051636 文本: 一般公共预算收入增长 5%
In [8]python · cell 19
python
print("新排名:")
for o in np.argsort(scores)[::-1]:
print(o+1)
print("分数和相应的文本:")
sorted_indices = np.argsort(scores)[::-1]
for rank, index in enumerate(sorted_indices, start=1):
# 打印排名和对应的分数
print(f"排名 {rank}, 分数: {scores[index]}")
# 打印对应的文档内容,这里假设 pairs[index][1] 是文档内容
print("文本:")
print(pairs[index][1])
print('') # 在文档之间添加空行以便区分Output
新排名: 3 4 1 2 8 7 10 5 6 9 分数和相应的文本: 排名 1, 分数: 8.211084365844727 文本: 人均地区生产总值、全 员劳动生产率、万元地区生产总值能耗水耗等多项指标保持全国省级地区最优水平 排名 2, 分数: 7.3925018310546875 文本: 提出今年全市经济社 会发展的主要预期目标是:地区生产总值增长 5%左右 排名 3, 分数: 7.281141757965088 文本: 全市地区生产总 值增长 5.2%、约 4.4 万亿元 排名 4, 分数: 7.037868499755859 文本: 数字经济增加值占地区生产总值比重达 42.9% 排名 5, 分数: 3.7360196113586426 文本: 一般公共预算收入增长 8.2%、突破 6000 亿元 排名 6, 分数: 3.665830612182617 文本: 推动经济实现质的有效提升和量的合理增长 排名 7, 分数: 1.5210455656051636 文本: 一般公共预算收入增长 5% 排名 8, 分数: 0.8092600107192993 文本: 居民收入增长与经济增长同步 排名 9, 分数: 0.8092600107192993 文本: 居民收入增长与经济增长同步 排名 10, 分数: -2.139439105987549 文本: 居民消费价格涨幅 3%左右
2.3 结合查询扩展的重排序
In [9]python · cell 21
python
# 接下来把之前获得的结果排序前5名传递给LLM
original_query = "推动北京财政收入增长因素是什么"
generated_queries = [
"什么推动了北京市的财政收入增长?",
"什么对北京市的财政收入增长作出贡献?",
"北京市财政收入增长依赖于什么?",
"什么有助于北京市的财政收入增长?",
"北京市财政收入增长受到哪些方面的影响?",
]In [10]python · cell 22
python
queries = [original_query] + generated_queries
results = chroma_collection.query(query_texts=queries, n_results=10, include=['documents', 'embeddings'])
retrieved_documents = results['documents']In [11]python · cell 23
python
# 删除检索文档中的重复数据
unique_documents = set()
for documents in retrieved_documents:
for document in documents:
unique_documents.add(document)
unique_documents = list(unique_documents)In [12]python · cell 24
python
# 再次创建pairs,可以增强查询的检索记过与原始查询的相关性
# 从中选择最佳的5个结果传递给LLM
pairs = []
for doc in unique_documents:
pairs.append([original_query, doc])In [13]python · cell 25
python
scores = cross_encoder.predict(pairs)In [14]python · cell 26
python
print("分数和对应的文本:")
# 打印每个文档的分数和内容
for score, (query, document) in zip(scores, pairs):
print(f"分数: {score}")
print("文本:")
print(word_wrap(document))
print('') # 在文档间添加空行以便区分Output
分数和对应的文本: 分数: 7.858731269836426 文本: 支持北京证券交易所扩容提质、上市公司数量增至开市时 的近三倍 分数: 7.336370944976807 文本: 是中共北京市委带领全市人民攻坚克难、艰苦奋斗的结果 分数: 7.265048503875732 文本: 促进北京普惠健 康保可持续发展 分数: 5.1238274574279785 文本: 为推进中国式现代化作出首都贡献 分数: 0.4882378578186035 文本: 一般公共预算收入增长 5% 分数: 1.943619728088379 文本: 持续增加城乡居民收入 分数: 0.3154720366001129 文本: 进一步优化提升首都 功能 分数: 3.9309439659118652 文本: 提高财政投入力度 分数: 4.740274429321289 文本: 更好发挥积极财政政策作用 分数: 5.981910705566406 文本: 深化健康北京建设 分数: 6.244244575500488 文本: 在中共北京市委坚强 领导下 分数: 3.023956060409546 文本: 首都文化持续繁荣发展 分数: 7.342596530914307 文本: 支持北京证券交易所深化改革和高质 量发展 分数: 5.825043678283691 文本: 我代表北京市人民政府 分数: 3.903216600418091 文本: 首都功能优化还有很大提升空间; 经济持续回升基础不牢固 分数: 0.8338263630867004 文本: 居民收入增长与经济增长同步 分数: 5.512838840484619 文本: 吸引国际组织和机构在京落地
In [15]python · cell 27
python
# 打印新排名和对应的结果
print("新排名:")
for o in np.argsort(scores)[::-1]:
print(o+1)
print("\n")
sorted_indices = np.argsort(scores)[::-1]
for rank, index in enumerate(sorted_indices, start=1):
# 打印排名和对应的分数
print(f"排名 {rank}, 分数: {scores[index]}")
# 打印对应的文档内容,这里假设 pairs[index][1] 是文档内容
print("文本:")
print(pairs[index][1])
print('') # 在文档之间添加空行以便区分Output
新排名: 1 13 2 3 11 10 14 17 4 9 8 15 12 6 16 5 7 排名 1, 分数: 7.858731269836426 文本: 支持北京证券交易所扩容提质、上市公司数量增至开市时 的近三倍 排名 2, 分数: 7.342596530914307 文本: 支持北京证券交易所深化改革和高质 量发展 排名 3, 分数: 7.336370944976807 文本: 是中共北京市委带领全市人民攻坚克难、艰苦奋斗的结果 排名 4, 分数: 7.265048503875732 文本: 促进北京普惠健 康保可持续发展 排名 5, 分数: 6.244244575500488 文本: 在中共北京市委坚强 领导下 排名 6, 分数: 5.981910705566406 文本: 深化健康北京建设 排名 7, 分数: 5.825043678283691 文本: 我代表北京市人民政府 排名 8, 分数: 5.512838840484619 文本: 吸引国际组织和机构在京落地 排名 9, 分数: 5.1238274574279785 文本: 为推进中国式现代化作出首都贡献 排名 10, 分数: 4.740274429321289 文本: 更好发挥积极财政政策作用 排名 11, 分数: 3.9309439659118652 文本: 提高财政投入力度 排名 12, 分数: 3.903216600418091 文本: 首都功能优化还有很大提升空间; 经济持续回升基础不牢固 排名 13, 分数: 3.023956060409546 文本: 首都文化持续繁荣发展 排名 14, 分数: 1.943619728088379 文本: 持续增加城乡居民收入 排名 15, 分数: 0.8338263630867004 文本: 居民收入增长与经济增长同步 排名 16, 分数: 0.4882378578186035 文本: 一般公共预算收入增长 5% 排名 17, 分数: 0.3154720366001129 文本: 进一步优化提升首都 功能
In [ ]python · cell 28
python
