使用RAG链与FewShotPromptTemplate时出现KeyError: 'context'求助
问题分析与解决方案
问题描述
用户尝试实现带Few-Shot示例的RAG链,希望Prompt包含示例和向量数据库返回的上下文,但运行时触发KeyError: 'context',原代码如下:
import pandas as pd import numpy as np from langchain.chains import RetrievalQA from langchain.embeddings import HuggingFaceEmbeddings from langchain.llms import HuggingFacePipeline from langchain.text_splitter import RecursiveCharacterTextSplitter, CharacterTextSplitter from langchain.vectorstores import FAISS from langchain.document_loaders import PyPDFDirectoryLoader from langchain.document_loaders import PyPDFLoader, Docx2txtLoader, TextLoader, DataFrameLoader import torch from transformers import AutoModelForCausalLM, AutoTokenizer,pipeline import os from langchain import PromptTemplate, FewShotPromptTemplate from langchain.schema.runnable import RunnablePassthrough model_name = 'AIMH/mental-longformer-base-4096' model_kwargs = {'device':'cuda'} encode_kwargs = {'normalize_embeddings':False} embedding= HuggingFaceEmbeddings( model_name = model_name, model_kwargs = model_kwargs, encode_kwargs = encode_kwargs ) document_path = "/content/drive/MyDrive/Colab_Notebooks/papers" indicators = ''' "An overwhelming sense that one can't escape their current situation or problems." "Alcohol or other substance use" "Disconnection from friends, family, and social activities." "Believing that nothing will ever get better or change." ''' # to df indicators = pd.DataFrame(indicators.split('\n'), columns=['indicators']) # load document loader = PyPDFDirectoryLoader(document_path) documents = loader.load() # make indicators a Document and append to document_splitted df_loader = DataFrameLoader(indicators, page_content_column="indicators") documents.extend(df_loader.load()) text_splitter = CharacterTextSplitter(chunk_size=500, chunk_overlap=10) chunked_documents = text_splitter.split_documents(documents) def create_db(document_splitted, embedding_model_instance): model_vectorstore = FAISS db=None try: content = [] metadata = [] for d in document_splitted: content.append(d.page_content) metadata.append({'source': d.metadata}) db=model_vectorstore.from_texts(content, embedding_model_instance, metadata) except Exception as error: print(error) return db db = create_db(chunked_documents, embedding) #store the db locally for future use db.save_local('db.index') retriever = db.as_retriever(search_type="similarity", search_kwargs={"k": 2}) model_path= "TheBloke/zephyr-7B-beta-AWQ" task = "text-generation" model_kwargs={ "temperature": 0, "max_length": 512, "do_sample": True, "top_k": 50, "top_p": 0.95, "num_return_sequences": 1 } pipeline_kwargs={ "repetition_penalty":1.1 } from awq import AutoAWQForCausalLM model = AutoAWQForCausalLM.from_quantized(model_path, fuse_layer=True,trust_remote_code = False, safetensors = True) tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code = False) pipe = pipeline( model=model, tokenizer=tokenizer, device="cuda", task=task ) llm = HuggingFacePipeline(pipeline = pipe, model_kwargs=model_kwargs, pipeline_kwargs=pipeline_kwargs) post1 = ''' my roomate drives me crazy. she bullies me and says horrible things. i am a very anxious person so i just hide in my room all day. i have not spoken to my family in weeks and have lost 10 pounds. what to do. ''' post1_label = 'severe' userid_1 = 1 post1_evidence = ["says horrible things", "I'm a very anxious person", "she bullies"] post2 = ''' I think i am depressed, i do feel like eating or going to the gym. what do i do? ''' post2_label = 'moderate' userid_2 = 1 examples = [ { "post": post1, "evidence": post1_evidence } ] example_template = """ {context} ###POST: {question} ###EVIDENCE: {evidence} """ prefix = """ You are an expert psychologist. You have a received information that the post's author is in one of 'Severe','Moderate',or 'Low' risk of depression. Use the following pieces of context to select the spans of text that provide evidence of the risk level. If you don't know the answer return an empty string (""). Do not make up an answer. """ suffix = """ {context} ###POST: {question} ###EVIDENCE: """ example_prompt = PromptTemplate( input_variables=["context","question", "evidence"], template=example_template ) few_shot_prompt_template = FewShotPromptTemplate( examples=examples, example_prompt=example_prompt, prefix=prefix, suffix=suffix, input_variables=["context","question"],#These variables are used in the prefix and suffix example_separator="\n\n" ) def gen_resp(retriever, question): rag_custom_prompt = few_shot_prompt_template context = "\n".join(doc.page_content for doc in retriever.get_relevant_documents(query = question)) rag_chain = ( {"context": lambda x: context, "question": RunnablePassthrough()} | rag_custom_prompt | llm ) answer = rag_chain.invoke(question) return answer gen_resp(retriever, post2)
运行后报错:
KeyError Traceback (most recent call last) <ipython-input-11-8208635e14b4> in <cell line: 25>() 23 return answer 24 ---> 25 gen_resp(retriever, post2) 9 frames /usr/local/lib/python3.10/dist-packages/langchain_core/prompts/few_shot.py in <dictcomp>(.0) 146 examples = self._get_examples(**kwargs) 147 examples = [ --> 148 {k: e[k] for k in self.example_prompt.input_variables} for e in examples 149 ] 150 # Format the examples. KeyError: 'context'
错误原因
- 示例字段不匹配:
example_prompt的input_variables定义了["context","question", "evidence"],但提供的examples列表中每个示例只有post和evidence字段,既没有context,也没有模板要求的question字段(用post代替了)。 - 上下文传递逻辑问题:链外提前获取的
context无法被FewShotPromptTemplate读取,模板渲染示例时会优先从examples中提取context字段。
修复步骤
步骤1:修正示例结构
给每个示例补充context字段,并将post重命名为question,和模板变量名对齐:
# 为示例生成对应的上下文 post1_context = "\n".join(doc.page_content for doc in retriever.get_relevant_documents(query=post1)) examples = [ { "question": post1, "evidence": post1_evidence, "context": post1_context } ]
步骤2:优化RAG链的上下文传递
将Retriever整合到Runnable链中,动态获取上下文,避免提前获取的冗余:
def gen_resp(retriever, question): rag_custom_prompt = few_shot_prompt_template rag_chain = ( {"context": retriever | (lambda docs: "\n".join(doc.page_content for doc in docs)), "question": RunnablePassthrough()} | rag_custom_prompt | llm ) answer = rag_chain.invoke(question) return answer
完整修改后代码
import pandas as pd import numpy as np from langchain.chains import RetrievalQA from langchain.embeddings import HuggingFaceEmbeddings from langchain.llms import HuggingFacePipeline from langchain.text_splitter import RecursiveCharacterTextSplitter, CharacterTextSplitter from langchain.vectorstores import FAISS from langchain.document_loaders import PyPDFDirectoryLoader from langchain.document_loaders import PyPDFLoader, Docx2txtLoader, TextLoader, DataFrameLoader import torch from transformers import AutoModelForCausalLM, AutoTokenizer,pipeline import os from langchain import PromptTemplate, FewShotPromptTemplate from langchain.schema.runnable import RunnablePassthrough model_name = 'AIMH/mental-longformer-base-4096' model_kwargs = {'device':'cuda'} encode_kwargs = {'normalize_embeddings':False} embedding= HuggingFaceEmbeddings( model_name = model_name, model_kwargs = model_kwargs, encode_kwargs = encode_kwargs ) document_path = "/content/drive/MyDrive/Colab_Notebooks/papers" indicators = ''' "An overwhelming sense that one can't escape their current situation or problems." "Alcohol or other substance use" "Disconnection from friends, family, and social activities." "Believing that nothing will ever get better or change." ''' # to df indicators = pd.DataFrame(indicators.split('\n'), columns=['indicators']) # load document loader = PyPDFDirectoryLoader(document_path) documents = loader.load() # make indicators a Document and append to document_splitted df_loader = DataFrameLoader(indicators, page_content_column="indicators") documents.extend(df_loader.load()) text_splitter = CharacterTextSplitter(chunk_size=500, chunk_overlap=10) chunked_documents = text_splitter.split_documents(documents) def create_db(document_splitted, embedding_model_instance): model_vectorstore = FAISS db=None try: content = [] metadata = [] for d in document_splitted: content.append(d.page_content) metadata.append({'source': d.metadata}) db=model_vectorstore.from_texts(content, embedding_model_instance, metadata) except Exception as error: print(error) return db db = create_db(chunked_documents, embedding) #store the db locally for future use db.save_local('db.index') retriever = db.as_retriever(search_type="similarity", search_kwargs={"k": 2}) model_path= "TheBloke/zephyr-7B-beta-AWQ" task = "text-generation" model_kwargs={ "temperature": 0, "max_length": 512, "do_sample": True, "top_k": 50, "top_p": 0.95, "num_return_sequences": 1 } pipeline_kwargs={ "repetition_penalty":1.1 } from awq import AutoAWQForCausalLM model = AutoAWQForCausalLM.from_quantized(model_path, fuse_layer=True,trust_remote_code = False, safetensors = True) tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code = False) pipe = pipeline( model=model, tokenizer=tokenizer, device="cuda", task=task ) llm = HuggingFacePipeline(pipeline = pipe, model_kwargs=model_kwargs, pipeline_kwargs=pipeline_kwargs) post1 = ''' my roomate drives me crazy. she bullies me and says horrible things. i am a very anxious person so i just hide in my room all day. i have not spoken to my family in weeks and have lost 10 pounds. what to do. ''' post1_label = 'severe' userid_1 = 1 post1_evidence = ["says horrible things", "I'm a very anxious person", "she bullies"] post2 = ''' I think i am depressed, i do feel like eating or going to the gym. what do i do? ''' post2_label = 'moderate' userid_2 = 1 # 修正示例结构:补充context字段,将post改为question post1_context = "\n".join(doc.page_content for doc in retriever.get_relevant_documents(query=post1)) examples = [ { "question": post1, "evidence": post1_evidence, "context": post1_context } ] example_template = """ {context} ###POST: {question} ###EVIDENCE: {evidence} """ prefix = """ You are an expert psychologist. You have a received information that the post's author is in one of 'Severe','Moderate',or 'Low' risk of depression. Use the following pieces of context to select the spans of text that provide evidence of the risk level. If you don't know the answer return an empty string (""). Do not make up an answer. """ suffix = """ {context} ###POST: {question} ###EVIDENCE: """ example_prompt = PromptTemplate( input_variables=["context","question", "evidence"], template=example_template ) few_shot_prompt_template = FewShotPromptTemplate( examples=examples, example_prompt=example_prompt, prefix=prefix, suffix=suffix, input_variables=["context","question"],#These variables are used in the prefix and suffix example_separator="\n\n" ) def gen_resp(retriever, question): rag_custom_prompt = few_shot_prompt_template # 优化上下文获取逻辑,整合到Runnable链中 rag_chain = ( {"context": retriever | (lambda docs: "\n".join(doc.page_content for doc in docs)), "question": RunnablePassthrough()} | rag_custom_prompt | llm ) answer = rag_chain.invoke(question) return answer gen_resp(retriever, post2)
内容的提问来源于stack exchange,提问作者laBouz
相关产品推荐
相关产品推荐

