30.RAG进阶(Advanced RAG)-后检索(PostRetrieval)-RAGFusion(RAG融合)
内容参考于:图灵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])

更多推荐


所有评论(0)