内容参考于:图灵AI大模型全栈

RAG-Fusion(RAG融合)

在多个查询检索后会产生很多的上下文,但并不是所有的上下文都与问题有关系,不相关的上下文可能出现在相关上下文的前面,这就会导致答案不准确

RAG-Fusion是一种搜索方式,它把生成多个问题的文档进行重排或进行融合

排序算法使用RRF,它的公式是1 / (rank+k),其中rank是当前排序的序号从0开始,k是一个常量随便设置

它的用法逻辑,如下,有Doc1、Doc2、Doc3、Doc4这4个文档,还有Question A、Question B、Question C、Question D4个问题,4个问题分别查出了Doc1、Doc2、Doc3、Doc4这4个文档,但是它们的排序的序号不同

Question A: Doc1 Doc4 Doc3 Doc2 Question B: Doc3 Doc1 Doc2 Doc4 Question C: Doc4 Doc3 Doc1 Doc2 Question D: Doc2 Doc1 Doc4 Doc3

Doc1的排序序号(rank): Question A rank: 0 Question B rank: 1 Question C rank: 2 Question D rank: 1

Doc2的排序序号(rank): Question A rank: 3 Question B rank: 2 Question C rank: 3 Question D rank: 0 Doc3的排序序号(rank): Question A rank: 2 Question B rank: 0 Question C rank: 1 Question D rank: 3 Doc4的排序序号(rank): Question A rank: 1 Question B rank: 3 Question C rank: 0 Question D rank: 2

然后现在k的值是60,如下方的结果

Doc1 Reciprocal Rank (Question A): 1 / (60 + 0) = 1 / 60 Reciprocal Rank (Question B): 1 / (60 + 1) = 1 / 61 Reciprocal Rank (Question C): 1 / (60 + 2) = 1 / 62 Reciprocal Rank (Question D): 1 / (60 + 1) = 1 / 61

Doc1用于重新排序的结果:0.0656

RRF(Doc1): 1 / 60 + 1 / 61 + 1 / 62 + 1 / 61 ≈ 0.0656 Doc2 Reciprocal Rank (Question A): 1 / (60 + 3) = 1 / 63 Reciprocal Rank (Question B): 1 / (60 + 2) = 1 / 62 Reciprocal Rank (Question C): 1 / (60 + 3) = 1 / 63 Reciprocal Rank (Question D): 1 / (60 + 0) = 1 / 60

Doc2用于重新排序的结果:0.0646

RRF(Doc2): 1 / 63 + 1 / 62 + 1 / 63 + 1 / 60 ≈ 0.0645 Doc3 Reciprocal Rank (Question A): 1 / (60 + 2) = 1 / 62 Reciprocal Rank (Question B): 1 / (60 + 0) = 1 / 60 Reciprocal Rank (Question C): 1 / (60 + 1) = 1 / 61 Reciprocal Rank (Question D): 1 / (60 + 3) = 1 / 63

Doc3用于重新排序的结果:0.0651

RRF(Doc3): 1 / 62 + 1 / 60 + 1 / 61 + 1 / 63 ≈ 0.0651 Doc4 Reciprocal Rank (Question A): 1 / (60 + 1) = 1 / 61 Reciprocal Rank (Question B): 1 / (60 + 3) = 1 / 63 Reciprocal Rank (Question C): 1 / (60 + 0) = 1 / 60 Reciprocal Rank (Question D): 1 / (60 + 2) = 1 / 62

Doc4用于重新排序的结果: 0.0651

RRF(Doc4): 1 / 61 + 1 / 63 + 1 / 60 + 1 / 62 ≈ 0.0651

然后现在的排序0.0656>0.0651>0.0645

Doc1是第一个,Doc4和Doc3是第二个(这里看具体实现也可能是第二和第三,谁是第二谁是第三看代码的具体实现),Doc2是第三个,这样就完成了重拍和融合(把A、B、C、D查出来的文档放到一起然后去重后生成一个新的),它有一个问题如果某个文档只查出了一次,就是说假设Doc5只在Question D中查出来了,那么Doc5的排序就会很低

这样排完序后再去TOPK个(也就是前x个)

代码

下方是代码运行打印的日志,可以看到input_variables是通过langsmith下拉的提示词,用来把问题生成多个,之前把问题生成多个是通过预索引库里提供的提示词来实现,这里是直接写提示词来实现

下拉的提示词:input_variables=['original_query'] input_types={} partial_variables={} metadata={'lc_hub_owner': 'langchain-ai', 'lc_hub_repo': 'rag-fusion-query-generation', 'lc_hub_commit_hash': '478b448e096b977446865108fad34282e6e1a84ae8b8540572ed0df238229a11'} messages=[SystemMessagePromptTemplate(prompt=PromptTemplate(input_variables=[], input_types={}, partial_variables={}, template='You are a helpful assistant that generates multiple search queries based on a single input query.'), additional_kwargs={}), HumanMessagePromptTemplate(prompt=PromptTemplate(input_variables=['original_query'], input_types={}, partial_variables={}, template='Generate multiple search queries related to: {original_query}'), additional_kwargs={}), HumanMessagePromptTemplate(prompt=PromptTemplate(input_variables=[], input_types={}, partial_variables={}, template='OUTPUT (4 queries):'), additional_kwargs={})]
--------------问题检索到的内容-------------
['1. 人工智能在医疗领域的应用案例与发展趋势  ', '2. 人工智能在金融行业中的实际应用场景(如风控、智能投顾)  ', '3. 工业制造中人工智能的应用:智能制造、预测性维护与质量检测  ', '4. 人工智能在教育、农业、交通等垂直行业的创新应用综述']
[Document(id='4777c270-e74f-46f8-b8a3-2d32c3168aff', metadata={}, page_content='人工智能在医疗诊断中的应用。'), Document(id='e0afae9a-9cda-43e3-b437-4b670b0bbe31', metadata={}, page_content='人工智能在制造业的应用。'), Document(id='2712bb3f-4c46-4b3c-91f1-deb486159597', metadata={}, page_content='人工智能如何影响未来就业市场。'), Document(id='2a6da101-9104-424d-ade4-28de2ce559f3', metadata={}, page_content='人工智能在金融风险管理中的应用。')]
[Document(id='2a6da101-9104-424d-ade4-28de2ce559f3', metadata={}, page_content='人工智能在金融风险管理中的应用。'), Document(id='2712bb3f-4c46-4b3c-91f1-deb486159597', metadata={}, page_content='人工智能如何影响未来就业市场。'), Document(id='e0afae9a-9cda-43e3-b437-4b670b0bbe31', metadata={}, page_content='人工智能在制造业的应用。'), Document(id='4777c270-e74f-46f8-b8a3-2d32c3168aff', metadata={}, page_content='人工智能在医疗诊断中的应用。')]
[Document(id='e0afae9a-9cda-43e3-b437-4b670b0bbe31', metadata={}, page_content='人工智能在制造业的应用。'), Document(id='4777c270-e74f-46f8-b8a3-2d32c3168aff', metadata={}, page_content='人工智能在医疗诊断中的应用。'), Document(id='2712bb3f-4c46-4b3c-91f1-deb486159597', metadata={}, page_content='人工智能如何影响未来就业市场。'), Document(id='07b9d4d5-770d-491a-b44a-596563bf7e5c', metadata={}, page_content='人工智能如何提升供应链效率。')]
[Document(id='e0afae9a-9cda-43e3-b437-4b670b0bbe31', metadata={}, page_content='人工智能在制造业的应用。'), Document(id='2712bb3f-4c46-4b3c-91f1-deb486159597', metadata={}, page_content='人工智能如何影响未来就业市场。'), Document(id='4777c270-e74f-46f8-b8a3-2d32c3168aff', metadata={}, page_content='人工智能在医疗诊断中的应用。'), Document(id='07b9d4d5-770d-491a-b44a-596563bf7e5c', metadata={}, page_content='人工智能如何提升供应链效率。')]
--------------问题融合后的内容-------------
D:\daimacunfangdi\PythonProject\03AdvancedRAG\03.RAG-Fusion.py:95: LangChainBetaWarning: The function `loads` is in beta. It is actively being worked on, so the API may change.
  (loads(doc), score)  for doc, score in sorted_Data
D:\daimacunfangdi\PythonProject\03AdvancedRAG\03.RAG-Fusion.py:95: LangChainPendingDeprecationWarning: The default value of `allowed_objects` will change in a future version. Pass an explicit list of allowed classes (or 'messages' for untrusted input that contains only chat messages) to suppress this warning.
  (loads(doc), score)  for doc, score in sorted_Data
[(Document(id='e0afae9a-9cda-43e3-b437-4b670b0bbe31', metadata={}, page_content='人工智能在制造业的应用。'), 0.06585580821434867), (Document(id='4777c270-e74f-46f8-b8a3-2d32c3168aff', metadata={}, page_content='人工智能在医疗诊断中的应用。'), 0.06532656778558418), (Document(id='2712bb3f-4c46-4b3c-91f1-deb486159597', metadata={}, page_content='人工智能如何影响未来就业市场。'), 0.06478053939714437), (Document(id='2a6da101-9104-424d-ade4-28de2ce559f3', metadata={}, page_content='人工智能在金融风险管理中的应用。'), 0.04841269841269841), (Document(id='07b9d4d5-770d-491a-b44a-596563bf7e5c', metadata={}, page_content='人工智能如何提升供应链效率。'), 0.015873015873015872)]
人工智能在制造业的应用。 0.06585580821434867
人工智能在医疗诊断中的应用。 0.06532656778558418
人工智能如何影响未来就业市场。 0.06478053939714437
人工智能在金融风险管理中的应用。 0.04841269841269841
人工智能如何提升供应链效率。 0.015873015873015872

下拉的提示词说明

# 整个提示词模板的输入变量:调用这个模板时,必须传入的参数
input_variables=['original_query']  
# 注释:只有一个参数 original_query,就是用户的原始问题(比如用户问"人工智能有哪些应用",就会把这句话填进模板里)

input_types={}  
# 注释:输入变量的类型声明,空字典代表用默认类型(字符串),一般不用改

partial_variables={}  
# 注释:预填充变量,也就是提前写死、不用每次调用都传的内容,这里没有提前固定的内容,所以为空

metadata={
    'lc_hub_owner': 'langchain-ai',
    'lc_hub_repo': 'rag-fusion-query-generation',
    'lc_hub_commit_hash': '478b448e096b977446865108fad34282e6e1a84ae8b8540572ed0df238229a11'
}
# 注释:元数据,记录这个提示词的来源信息
# - lc_hub_owner:这个提示词的作者/组织,这里是 LangChain 官方
# - lc_hub_repo:提示词所在的仓库名,专门用来做 RAG 融合的查询生成
# - lc_hub_commit_hash:代码提交的哈希值,相当于这个提示词的「版本号」,用来定位具体是哪个版本

messages=[
    # 第1条消息:系统提示词,给大模型定身份和核心任务
    SystemMessagePromptTemplate(
        prompt=PromptTemplate(
            input_variables=[],  # 这条系统消息不需要传任何参数,内容是固定的
            input_types={},
            partial_variables={},
            template='You are a helpful assistant that generates multiple search queries based on a single input query.'
            # 注释:系统指令原文:你是一个助手,要基于单个输入问题,生成多个搜索查询
        ),
        additional_kwargs={}  # 额外参数,一般用来传大模型的特殊配置,这里为空
    ),

    # 第2条消息:用户提问模板,把用户的原始问题填进去
    HumanMessagePromptTemplate(
        prompt=PromptTemplate(
            input_variables=['original_query'],  # 需要传入用户原始问题
            input_types={},
            partial_variables={},
            template='Generate multiple search queries related to: {original_query}'
            # 注释:模板内容,{original_query} 是占位符,会被替换成用户的真实问题
        ),
        additional_kwargs={}
    ),

    # 第3条消息:格式约束,强制大模型按要求输出
    HumanMessagePromptTemplate(
        prompt=PromptTemplate(
            input_variables=[],  # 内容固定,不需要传参
            input_types={},
            partial_variables={},
            template='OUTPUT (4 queries):'
            # 注释:告诉大模型「输出4个查询」,用这句话开头,约束输出的数量和格式
        ),
        additional_kwargs={}
    )
]

生成的多个问题

[
    '1. 人工智能在医疗领域的应用案例与发展趋势  ',
    '2. 人工智能在金融行业中的实际应用场景(如风控、智能投顾)  ',
    '3. 工业制造中人工智能的应用:智能制造、预测性维护与质量检测  ',
    '4. 人工智能在教育、农业、交通等垂直行业的创新应用综述'
]

文档说明

Document(
    id='4777c270-e74f-46f8-b8a3-2d32c3168aff',  # 文档唯一ID
    metadata={},  # 文档元数据
    page_content='人工智能在医疗诊断中的应用。'  # 文档正文内容
)

融合重拍后

[
    # 每个元素是一个元组:(Document对象, 融合得分),按得分从高到低排序
    (Document(id='e0afae9a...', page_content='人工智能在制造业的应用。'), 0.06585580821434867),
    (Document(id='4777c270...', page_content='人工智能在医疗诊断中的应用。'), 0.06532656778558418),
    (Document(id='2712bb3f...', page_content='人工智能如何影响未来就业市场。'), 0.06478053939714437),
    (Document(id='2a6da101...', page_content='人工智能在金融风险管理中的应用。'), 0.04841269841269841),
    (Document(id='07b9d4d5...', page_content='人工智能如何提升供应链效率。'), 0.015873015873015872)
]

代码

# 导入操作系统模块:读取环境变量,比如API密钥、LangSmith配置
import os

# 导入Chroma向量数据库:轻量级本地向量库,存储文本向量,执行语义相似度搜索
from langchain_chroma import Chroma
# 导入字符串输出解析器:把大模型返回的消息对象转成纯文本字符串
from langchain_core.output_parsers import StrOutputParser
# 导入HuggingFace嵌入模型:加载本地嵌入模型,把文本转换成数字向量
from langchain_huggingface import HuggingFaceEmbeddings
# 导入OpenAI格式大模型客户端:通义千问兼容OpenAI接口,可直接复用
from langchain_openai import ChatOpenAI

# 导入LangSmith客户端:用来连接LangChain官方的提示词仓库(Hub),拉取现成的优化好的提示词模板
from langsmith import Client
# 导入LangChain的序列化/反序列化工具:dumps把Document对象转成字符串,loads把字符串转回Document
# 为什么需要:后面融合分数的时候,要用文档当字典的key,但Python对象不能直接当key,转成字符串才行
from langchain_core.load import dumps, loads

# 导入环境变量加载工具:读取项目根目录的.env文件,加载配置
from dotenv import load_dotenv

# 加载.env文件中的环境变量(包括大模型API密钥、LangSmith API密钥等)
# 为什么调用:敏感信息不硬写在代码里,更安全,切换环境也不用改代码
load_dotenv()


# -------------------------- 1. 准备测试知识库文本 --------------------------
# 模拟知识库的10条文本,覆盖不同主题,用来测试检索融合效果
# 入参来源:手动编写的测试数据,实际项目中就是从文档加载、切分得到的文本块
texts = [
    "人工智能在医疗诊断中的应用。",
    "人工智能如何提升供应链效率。",
    "NBA季后赛最新赛况分析。",
    "传统法式烘焙的五大技巧。",
    "红楼梦人物关系图谱分析。",
    "人工智能在金融风险管理中的应用。",
    "人工智能如何影响未来就业市场。",
    "人工智能在制造业的应用。",
    "今天天气怎么样",
    "人工智能伦理:公平性与透明度。"
]


# -------------------------- 2. 初始化嵌入模型与向量库 --------------------------
# 本地嵌入模型的文件路径(BGE中文大模型,专门做文本向量化)
# 入参来源:手动填写本地模型的绝对路径,需要提前下载好模型文件
embedding_model_path = r'D:\huanjing\ai模型\BAAI\bge-large-zh-v1___5'

# 初始化嵌入模型实例
# 为什么调用:向量数据库只能存向量、算相似度,需要嵌入模型完成「文本 → 数字向量」的转换
# 入参:model_name 指定模型的本地路径
embeddings_model = HuggingFaceEmbeddings(
    model_name=embedding_model_path
)

# 一步完成「向量库创建 + 文本入库 + 批量向量化」
# 为什么用from_texts:直接从纯文本列表创建向量库,不用手动转成Document对象,适合简单场景
# 入参说明:
#   texts=texts:要入库的纯文本列表
#       → 入参来源:上面定义的测试文本列表
#   embedding=embeddings_model:嵌入模型实例
#       → 入参来源:上面初始化好的embeddings_model
vectorstore = Chroma.from_texts(
    texts=texts, embedding=embeddings_model
)

# 把向量库转成标准检索器对象
# 为什么调用:统一检索接口,后续可以直接调用invoke检索,也可以配合map做批量检索
retriever = vectorstore.as_retriever()


# -------------------------- 3. 从LangSmith Hub拉取官方提示词模板 --------------------------
# 创建LangSmith客户端实例
# 为什么调用:用来连接LangChain的提示词仓库,拉取官方/社区优化好的提示词,不用自己从零写
# 注意:使用前需要在.env里配置 LANGCHAIN_API_KEY,否则无法拉取
client = Client()

# 从LangChain Hub拉取RAG Fusion专用的查询生成提示词模板
# 提示词仓库地址:https://smith.langchain.com/hub/langchain-ai/rag-fusion-query-generation
# 这个提示词是官方优化好的,专门用来「根据原始问题生成多个不同表述的相关查询」
# 入参说明:
#   第1个参数:提示词的仓库路径,固定写法,对应Hub上的模板
#   dangerously_pull_public_prompt=True:允许拉取公共提示词的开关
#       → 为什么要加:公共提示词默认需要显式确认才能拉取,防止误拉不受信任的模板
# 返回值:ChatPromptTemplate对象,也就是现成的提示词模板
prompt = client.pull_prompt(
    "langchain-ai/rag-fusion-query-generation",
    dangerously_pull_public_prompt=True
)
# 打印拉取到的提示词,方便查看官方模板写了什么
print(f'下拉的提示词:{prompt}')


# -------------------------- 4. 初始化大模型 --------------------------
# 初始化通义千问大模型(通过OpenAI兼容接口调用)
# 入参说明:
#   model:模型名称,qwen-plus是通义的主力通用模型
#   api_key:API密钥,从环境变量读取
#   base_url:API接口地址,从环境变量读取
llm = ChatOpenAI(
    model="qwen-plus",
    api_key=os.getenv("DASHSCOPE_API_KEY"),
    base_url=os.getenv("DASHSCOPE_BASE_URL")
)


# -------------------------- 5. 构建「多查询生成链」 --------------------------
# 创建多重查询生成链:输入原始问题 → 生成多个不同表述的查询问题列表
# 链式执行流程:
# 1. prompt:接收原始问题original_query,填充进提示词模板
# 2. llm:调用大模型,生成多个不同表述的查询,每行一个问题
# 3. StrOutputParser():把大模型返回的消息对象转成纯文本字符串
# 4. lambda x: x.split("\n"):按换行符切割字符串,把多行问题转成字符串列表
generate_queries = (
        prompt | llm | StrOutputParser() | (lambda x: x.split("\n"))
)

# 定义原始查询问题
original_query = "人工智能的应用"
# 调用链,生成多个扩展查询
# 入参:字典,key为original_query,value是原始问题
# 返回值:字符串列表,每个元素是一个生成的扩展查询
queries = generate_queries.invoke({"original_query": original_query})

print('--------------问题检索到的内容-------------')
# 打印生成的所有扩展查询,看看大模型生成了哪些不同问法
print(queries)

# 循环每个查询,分别执行检索,打印各自的结果
# 为什么要单独打印:直观看到每个查询搜到的文档不一样、排名也不一样
# 这就是为什么需要融合重排——单独一个查询的结果有局限性,多个结果融合更准
for i in queries:
    print(retriever.invoke(i))


# -------------------------- 6. 核心算法:RRF 倒数排名融合 --------------------------
def reciprocal_rank_fusion(results: list[list], k=60):
    """
    互逆排名融合算法(Reciprocal Rank Fusion, RRF)
    核心作用:把多个检索结果列表,按排名位置计算分数,合并后重新排序
    解决的问题:多个查询返回的文档列表,排名不一样,怎么合并出一个最优的排序
    核心思想:一个文档在越多列表里出现、且排名越靠前,最终得分就越高
    
    Args:
        results: 二维列表,每个元素是一个查询的检索结果列表(文档对象列表)
            比如 [[doc1, doc2], [doc2, doc3], [doc1, doc4]]
        k: 平滑参数,默认60。值越小,排名靠前的文档分数优势越大;值越大,排名影响越平缓
    Returns:
        列表,每个元素是(文档对象, 融合分数)的元组,按分数从高到低排序
    """

    # 初始化融合分数字典
    # key:序列化后的文档字符串(用来唯一标识一个文档)
    # value:该文档的累计融合分数
    fused_scores = {}

    # 遍历每一个查询的检索结果列表
    for docs in results:
        # 遍历当前列表里的每个文档,rank是文档的排名(从0开始,0是第1名)
        for rank, doc in enumerate(docs):
            # 把Document对象序列化成字符串
            # 为什么要序列化:Python的对象不能直接当字典的key,转成字符串才能存进fused_scores
            doc_str = dumps(doc)
            
            # 如果这个文档是第一次出现,初始化分数为0
            if doc_str not in fused_scores:
                fused_scores[doc_str] = 0
            
            # 计算当前排名的RRF分数,累加到总分里
            # 公式:1 / (排名 + k)
            # 排名越靠前(rank越小),分数越高;排名越靠后,分数越低
            # k是平滑系数,避免第1名的分数和后面差距过大
            fused_scores[doc_str] += 1 / (rank + k)

    # 按融合分数从高到低排序
    # sorted返回的是列表,每个元素是(文档字符串, 分数)的元组
    sorted_Data = sorted(fused_scores.items(), key=lambda x: x[1], reverse=True)
    
    # 把序列化的文档字符串,转回原来的Document对象
    # 最终结果格式:[(文档1, 分数1), (文档2, 分数2), ...],按分数降序排列
    reranked_results = [
        (loads(doc), score)  for doc, score in sorted_Data
    ]

    return reranked_results


# -------------------------- 7. 完整RAG Fusion链路 --------------------------
print('--------------问题融合后的内容-------------')
# 重新定义原始查询(和前面一致,保证逻辑连贯)
original_query = "人工智能的应用"

# 构建完整的RAG Fusion检索链
# 执行流程:
# 1. generate_queries:输入原始问题 → 生成多个扩展查询列表
# 2. retriever.map():批量检索,给每个查询都调用一次retriever.invoke,返回二维结果列表
#    → 为什么用map():retriever本身只能接收一个查询,map()把它变成能批量处理列表的能力
#    → 输入是查询列表,输出是结果的二维列表,正好匹配RRF函数的入参格式
# 3. reciprocal_rank_fusion:把多个检索结果用RRF算法融合重排
chain = generate_queries | retriever.map() | reciprocal_rank_fusion

# 调用完整链路,传入原始问题,得到最终融合重排后的结果
# 入参:字典,key为original_query,value是原始问题
# 返回值:列表,每个元素是(Document对象, 分数)的元组,按分数从高到低排序
result_list = chain.invoke({"original_query": original_query})
# 打印完整结果
print(result_list)

# 遍历结果,单独打印每个文档的内容和对应的融合分数
# 更直观地看到排序结果和得分差异
for i in result_list:
    print(i[0].page_content, i[1])


img

Logo

脑启社区是一个专注类脑智能领域的开发者社区。欢迎加入社区,共建类脑智能生态。社区为开发者提供了丰富的开源类脑工具软件、类脑算法模型及数据集、类脑知识库、类脑技术培训课程以及类脑应用案例等资源。

更多推荐