You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.01 20:54:57