如何从Chainlit Prompt正确获取input_variables?Llama2部署报错求助
问题修复方案
错误1:StuffDocumentsChain的context变量未找到
原因
自定义prompt模板中,Context: {}的占位符未正确对应input_variables里的context参数,导致LangChain无法识别该变量,进而报错document_variable_name context was not found in llm_chain input_variables: ['', 'question']。
修复代码
修改custom_prompt_template中的Context行:
custom_prompt_template = """Use the following pieces of information to answer the user's question. If you don't know the answer, please just say that you don't know the answer, don't try to make up an answer. Context: {context} Question: {question} Only returns the helpful answer below and nothing else. Helpful answer: """
错误2:UserSession.set()缺少必填参数
原因
在@cl.on_message回调中,错误使用cl.user_session.set("chain")获取会话中的chain对象。set()方法需要传入键和值两个参数,而获取对象应该使用get()方法;同时chain.acall()需要传入用户输入的文本内容,而非message对象本身。
修复代码
修改@cl.on_message中的逻辑:
@cl.on_message async def main(message): chain = cl.user_session.get("chain") cb = cl.AsyncLangchainCallbackHandler( stream_final_answer=True, answer_prefix_tokens=["FINAL", "ANSWER"] ) cb.answer_reached = True res = await chain.acall(message.content, callbacks=[cb]) answer = res["result"] sources = res["source_documents"] if sources: answer += "\nSources:\n" for idx, doc in enumerate(sources, 1): answer += f"{idx}. {doc.page_content[:150]}...\n" else: answer += "\nNo Sources Found" await cl.Message(content=answer).send()
完整修复后的代码
from langchain.prompts import PromptTemplate from langchain.embeddings import HuggingFaceEmbeddings from langchain.vectorstores.faiss import FAISS from langchain.llms import CTransformers from langchain.chains import RetrievalQA import chainlit as cl DB_FAISS_PATH = "vectorstores/db_faiss" custom_prompt_template = """Use the following pieces of information to answer the user's question. If you don't know the answer, please just say that you don't know the answer, don't try to make up an answer. Context: {context} Question: {question} Only returns the helpful answer below and nothing else. Helpful answer: """ def set_custom_prompt(): prompt = PromptTemplate(template=custom_prompt_template, input_variables=['context','question']) return prompt def load_llm(): print("*** Start the load_llm.") llm = CTransformers( model="llama-2-7b-chat.ggmlv3.q8_0.bin", model_type="llama", max_new_tokens=512, temperature=0.5 ) print("****** Finish the load_llm") return llm def retrieval_qa_chain(llm, prompt, db): qa_chain = RetrievalQA.from_chain_type( llm=llm, chain_type="stuff", retriever=db.as_retriever(search_kwargs={'k':2}), return_source_documents=True, chain_type_kwargs={'prompt': prompt} ) return qa_chain def qa_bot(): embeddings = HuggingFaceEmbeddings(model_name='sentence-transformers/all-MiniLM-L6-v2', model_kwargs={'device': 'cpu'}) db = FAISS.load_local(DB_FAISS_PATH, embeddings) print("*******FAISS.load_local() works well.") llm = load_llm() print("****** llm step works well.") qa_prompt = set_custom_prompt() qa = retrieval_qa_chain(llm, qa_prompt, db) print("******qa step works well.") return qa # chainlit #### @cl.on_chat_start async def start(): chain = qa_bot() msg = cl.Message(content="Starting the bot......") await msg.send() msg.content = "Hi, Welcome to the Medical Bot. What is your query?" await msg.update() cl.user_session.set('chain', chain) @cl.on_message async def main(message): chain = cl.user_session.get("chain") cb = cl.AsyncLangchainCallbackHandler( stream_final_answer=True, answer_prefix_tokens=["FINAL", "ANSWER"] ) cb.answer_reached = True res = await chain.acall(message.content, callbacks=[cb]) answer = res["result"] sources = res["source_documents"] if sources: answer += "\nSources:\n" for idx, doc in enumerate(sources, 1): answer += f"{idx}. {doc.page_content[:150]}...\n" else: answer += "\nNo Sources Found" await cl.Message(content=answer).send()
内容的提问来源于stack exchange,提问作者LeiDing
相关产品推荐
相关产品推荐

