无需嵌入与向量存储时,ConversationalRetrievalChain的retriever参数传什么?
问题
我正在开发一款可根据用户问题生成Oracle查询语句的聊天机器人,此前已通过SystemMessage传入表结构等数据库相关信息,HumanMessage传入用户问题,运行正常且能生成查询语句。现在想将其改造为交互式对话机器人,可根据用户需求生成/修改响应,但因为无需文档嵌入与向量存储,请问ConversationalRetrievalChain函数的retriever参数应该传入什么?
附上我的代码:
import logging import json import langchain import os from azure.identity import ClientSecretCredential from langchain.chains.conversational_retrieval.base import ConversationalRetrievalChain from langchain.prompts import SystemMessagePromptTemplate, HumanMessagePromptTemplate, ChatPromptTemplate,MessagesPlaceholder from langchain_openai import AzureChatOpenAI from langchain.memory import ConversationBufferMemory from src.dao.dbConnect import getFromDb langchain.verbose = False langchain.debug = False langchain.llm_cache = False logging.basicConfig(level=logging.INFO) logger = logging.getLogger("uvicorn") def getResponseFromModel(guide, question, environment): chat_history = [] memory = ConversationBufferMemory(memory_key="chat_history", return_messages=True) messages = [ SystemMessagePromptTemplate.from_template(guide), MessagesPlaceholder(variable_name="chat_history"), HumanMessagePromptTemplate.from_template(question) ] client_id = os.environ["CLIENT_ID"] client_key =os.environ["CLIENT_KEY"] credential = ClientSecretCredential( client_id=client_id, client_secret=client_key, tenant_id=os.environ["TENANT_ID"] ) token = credential.get_token("https://cognitiveservices.azure.com/.default") model = AzureChatOpenAI( azure_endpoint=os.environ["AZURE_OPENAI_ENDPOINT"], openai_api_version=os.environ["AZURE_OPENAI_API_VERSION"], azure_deployment=os.environ["AZURE_OPENAI_CHAT_DEPLOYMENT_NAME"], openai_api_key=token.token, temperature=0, top_p = 0.1 ) prompt = ChatPromptTemplate.from_messages(messages=messages) logger.info("Fetching Response from model") qa = ConversationalRetrievalChain.from_llm( llm= model, retriever= , #Here I am getting the issue memory=memory, combine_docs_chain_kwargs={"prompt": prompt} ) response = qa({"question": question, "chat_history": chat_history}) #response = model.invoke(message) response_content = json.loads(response["answer"].strip('```json').strip('```').strip()) result = "" if response_content['question_type'] == 'sql': df = getFromDb( str(response_content['response']).replace("sql", "").replace("`", ""), environment ) result = json.dumps({"query": (response_content['response']).replace("sql", "").replace("`", ""), "result": df}) else: result = json.dumps({"result": (response_content['response']).replace("sql", "").replace("`", "")}) return result
解答
首先明确:ConversationalRetrievalChain的核心设计目标是结合对话记忆与文档检索,如果你的场景完全不需要文档检索(无向量存储或外部文档查询需求),这个链并非最优选择,更推荐直接使用带对话记忆的ConversationChain或自定义对话链。
但如果一定要继续使用ConversationalRetrievalChain,可以传入一个返回空文档列表的假检索器,因为链内部需要调用retriever的get_relevant_documents方法,只要该方法返回空列表即可。具体有两种实现方式:
方式1:自定义DummyRetriever
实现一个简单的Retriever类,让其get_relevant_documents方法返回空列表:
from langchain.schema import BaseRetriever class DummyRetriever(BaseRetriever): def _get_relevant_documents(self, query): return [] # 创建链时使用 qa = ConversationalRetrievalChain.from_llm( llm=model, retriever=DummyRetriever(), memory=memory, combine_docs_chain_kwargs={"prompt": prompt} )
方式2:使用空向量存储生成Retriever(不推荐)
这种方式冗余,但如果一定要用,可以创建空的FAISS向量存储再生成retriever:
from langchain.vectorstores import FAISS from langchain.embeddings import FakeEmbeddings # 创建空向量存储 empty_vectorstore = FAISS.from_texts([""], FakeEmbeddings()) retriever = empty_vectorstore.as_retriever() # 传入链中 qa = ConversationalRetrievalChain.from_llm( llm=model, retriever=retriever, memory=memory, combine_docs_chain_kwargs={"prompt": prompt} )
更优方案:替换为ConversationChain
既然不需要检索,直接用ConversationChain更简洁,完全不需要retriever参数:
from langchain.chains import ConversationChain # 构建prompt时保留对话历史占位符 prompt = ChatPromptTemplate.from_messages(messages=messages) # 创建ConversationChain conversation_chain = ConversationChain( llm=model, memory=memory, prompt=prompt ) # 调用链获取响应 response = conversation_chain.invoke({"input": question}) # 后续处理逻辑不变,注意响应键为"response"而非"answer" response_content = json.loads(response["response"].strip('```json').strip('```').strip())
内容的提问来源于stack exchange,提问作者lakshya saxena
相关产品推荐
相关产品推荐

