在SageMaker部署Dolly2模型生成Embeddings时遇400错误求助
问题解决:SageMaker部署Dolly2后Embeddings生成报错400
错误根源
HF_TASK环境变量未在SageMaker推理容器中配置,本地代码设置该变量仅作用于本地进程,无法传递到云端容器。- 自定义的
inference.py缺少SageMaker推理必需的请求处理函数(input_fn/predict_fn/output_fn),无法正确解析请求并返回结果。
解决方案
1. 配置SageMaker容器环境变量
在创建模型或部署端点时,必须将HF_TASK添加到容器环境变量中:
- 代码部署方式:
from sagemaker.huggingface import HuggingFaceModel model = HuggingFaceModel( image_uri=your_hf_image_uri, model_data=your_model_s3_uri, role=your_sagemaker_execution_role, env={ "HF_TASK": "feature-extraction" # 关键配置 } ) model.deploy(instance_type="ml.g5.xlarge", initial_instance_count=1, endpoint_name="your-endpoint-name") - 控制台部署方式:在端点部署步骤的「环境变量」模块,添加
HF_TASK=feature-extraction键值对。
2. 修复inference.py推理脚本
补充必需的推理处理函数,同时支持文本生成、Embeddings、QA三类任务:
import torch from transformers import pipeline import json def model_fn(model_dir): # 初始化文本生成 pipeline text_gen_pipeline = pipeline( "text-generation", model=model_dir, torch_dtype=torch.bfloat16, trust_remote_code=True, device_map="auto", model_kwargs={"load_in_8bit": True}, ) tokenizer = text_gen_pipeline.tokenizer embeddings_model = text_gen_pipeline.model # 封装模型相关对象统一返回 return { "text_gen_pipeline": text_gen_pipeline, "tokenizer": tokenizer, "embeddings_model": embeddings_model } def input_fn(request_body, request_content_type): # 解析JSON格式输入 if request_content_type == "application/json": return json.loads(request_body) raise ValueError(f"不支持的内容类型: {request_content_type}") def predict_fn(input_data, model_objects): text_gen_pipeline = model_objects["text_gen_pipeline"] tokenizer = model_objects["tokenizer"] embeddings_model = model_objects["embeddings_model"] # 根据输入格式判断任务类型 if "inputs" in input_data: # 处理Embeddings生成请求 inputs = tokenizer( input_data["inputs"], truncation=True, padding="longest", return_tensors="pt" ).to("cuda" if torch.cuda.is_available() else "cpu") with torch.no_grad(): outputs = embeddings_model(**inputs) embeddings = outputs.last_hidden_state.mean(dim=1).squeeze(0).tolist() return {"embeddings": [embeddings]} elif "question" in input_data and "context" in input_data: # 处理QA请求 qa_outputs = text_gen_pipeline(input_data["question"], input_data["context"]) return {"qa_result": qa_outputs} elif "prompt" in input_data: # 处理文本生成请求 gen_outputs = text_gen_pipeline(input_data["prompt"]) return {"generated_text": gen_outputs} raise ValueError("无效输入格式。支持格式:{'inputs': '文本'}(Embeddings)、{'question': '问题', 'context': '上下文'}(QA)、{'prompt': '提示词'}(文本生成)") def output_fn(prediction, response_content_type): # 格式化JSON输出 if response_content_type == "application/json": return json.dumps(prediction) raise ValueError(f"不支持的响应内容类型: {response_content_type}")
3. 调整端点调用代码
移除本地无用的HF_TASK设置,直接发送对应格式的请求:
import json import boto3 def invoke_sagemaker_endpoint(): sagemaker_client = boto3.client("sagemaker-runtime") endpoint_name = 'XXX' # 替换为你的端点名称 # Embeddings请求 payload payload = {"inputs": "This is a large document."} response = sagemaker_client.invoke_endpoint( EndpointName=endpoint_name, ContentType="application/json", Body=json.dumps(payload), ) response_body = response["Body"].read().decode("utf-8") response_json = json.loads(response_body) if "embeddings" in response_json: return response_json["embeddings"][0] return None if __name__ == "__main__": embeddings_vector = invoke_sagemaker_endpoint() if embeddings_vector: print(embeddings_vector) else: print("响应中未找到Embeddings数据。")
4. 重新部署端点
将修改后的inference.py打包,配合正确的环境变量配置,重新创建模型并部署SageMaker端点。
内容的提问来源于stack exchange,提问作者Arpel
相关产品推荐
相关产品推荐

