Sagemaker+LangChain部署Llama2调用InvokeEndpoint遇ValueError问题
使用SageMaker+LangChain部署Llama 2时的请求格式错误解决
问题描述
尝试通过SageMaker和LangChain部署Llama 2模型做文本生成推理,编写问答链代码后运行触发以下错误:
ValueError: Error raised by inference endpoint: An error occurred (ModelError) when calling the InvokeEndpoint operation: Received client error (422) from primary with message "Failed to deserialize the JSON body into the target type: missing field `inputs` at line 1 column 966".
核心代码如下:
from langchain.docstore.document import Document example_doc_1 = """ Peter and Elizabeth took a taxi to attend the night party in the city. While in the party, Elizabeth collapsed and was rushed to the hospital. Since she was diagnosed with a brain injury, the doctor told Peter to stay besides her until she gets well. Therefore, Peter stayed with her at the hospital for 3 days without leaving. """ docs = [ Document( page_content=example_doc_1, ) ] from typing import Dict from langchain import PromptTemplate, SagemakerEndpoint from langchain.llms.sagemaker_endpoint import LLMContentHandler from langchain.chains.question_answering import load_qa_chain import json query = """How long was Elizabeth hospitalized? """ prompt_template = """Use the following pieces of context to answer the question at the end. {context} Question: {question} Answer:""" PROMPT = PromptTemplate( template=prompt_template, input_variables=["context", "question"] ) class ContentHandler(LLMContentHandler): content_type = "application/json" accepts = "application/json" def transform_input(self, prompt: str, model_kwargs: Dict) -> bytes: input_str = json.dumps({prompt: prompt, **model_kwargs}) return input_str.encode("utf-8") def transform_output(self, output: bytes) -> str: response_json = json.loads(output.read().decode("utf-8")) return response_json[0]["generated_text"] content_handler = ContentHandler() chain = load_qa_chain( llm=SagemakerEndpoint( endpoint_name="XYZ", credentials_profile_name="XYZ", region_name="XYZ", model_kwargs={"temperature": 1e-10}, content_handler=content_handler, ), prompt=PROMPT, ) chain({"input_documents": docs, "question": query}, return_only_outputs=True)
问题解答
1. 错误原因与修复方法
错误核心是SageMaker上的Llama 2推理容器要求请求体必须以inputs作为输入文本的键,但代码中ContentHandler的transform_input方法错误地将prompt内容本身作为键名,导致请求格式不符合模型解析要求,触发422参数错误。
修复只需修改transform_input方法,将请求体的键改为inputs:
def transform_input(self, prompt: str, model_kwargs: Dict) -> bytes: # 替换原错误的键名,使用模型要求的"inputs"字段 input_str = json.dumps({"inputs": prompt, **model_kwargs}) return input_str.encode("utf-8")
另外需注意,不同版本的Llama 2推理容器返回格式可能有差异,若修复后仍出现解析错误,可根据实际返回结构调整transform_output,比如部分容器直接返回generated_text字段:
def transform_output(self, output: bytes) -> str: response_json = json.loads(output.read().decode("utf-8")) return response_json.get("generated_text", "")
2. 代码中缺失的配置
代码主要缺失对Llama 2推理容器请求/响应格式的适配:
- 未遵循模型要求使用
inputs字段传递输入文本,构造了无效的请求键 - 未根据实际部署的Llama 2容器返回结构,调整输出解析逻辑
此外需确认SageMaker端点部署的是官方兼容的Llama 2推理镜像,若使用自定义镜像,需对应调整ContentHandler的请求和响应处理逻辑。
内容的提问来源于stack exchange,提问作者Aun Zaidi
相关产品推荐
相关产品推荐

