在Amazon SageMaker调用LLaVA v1.6 Mistral模型推理遇错求助
一、解决ModelError问题
1. 升级Transformers及依赖版本
LLaVA v1.6的llava_next架构需要较新版本的Transformers支持,建议在SageMaker部署时指定依赖版本:
创建requirements.txt文件,内容如下:
transformers>=4.40.0 accelerate>=0.29.0 pillow>=10.0.0 torch>=2.1.0
部署模型时,在HuggingFaceModel构造函数中通过requirements_file参数引入该文件,确保容器安装对应版本依赖。
2. 修正推理脚本的模型加载逻辑
避免使用Auto类自动加载,改用LLaVA-Next专用类,确保架构识别正确:
from transformers import LlavaNextProcessor, LlavaNextForConditionalGeneration import torch def model_fn(model_dir): processor = LlavaNextProcessor.from_pretrained(model_dir) model = LlavaNextForConditionalGeneration.from_pretrained( model_dir, torch_dtype=torch.float16, device_map="auto" ) return processor, model def predict_fn(data, model_and_processor): processor, model = model_and_processor # 处理base64格式的图片(如果Lambda传过来的是base64) if isinstance(data["image"], str): from PIL import Image import base64 from io import BytesIO image_bytes = base64.b64decode(data["image"]) image = Image.open(BytesIO(image_bytes)).convert("RGB") else: image = data["image"] prompt = data["prompt"] inputs = processor(prompt, image, return_tensors="pt").to(model.device) outputs = model.generate(**inputs, max_new_tokens=512) response = processor.decode(outputs[0], skip_special_tokens=True) return {"response": response}
3. 验证模型检查点完整性
确认S3中的模型文件完整,特别是config.json内model_type字段为llava_next,无文件损坏。可先在本地用相同依赖版本加载模型,验证正常后再部署到SageMaker。
二、Lambda + boto3调用SageMaker端点的正确方式
1. Lambda函数核心代码
确保Lambda执行角色拥有sagemaker:InvokeEndpoint权限,代码示例:
import boto3 import base64 import json def lambda_handler(event, context): # 解析API Gateway传入的请求体 request_body = json.loads(event["body"]) image_base64 = request_body["image"] prompt = request_body["prompt"] # 调用SageMaker端点 sagemaker_runtime = boto3.client("sagemaker-runtime") endpoint_name = "your-llava-endpoint-name" # 替换为你的端点名称 payload = json.dumps({ "image": image_base64, "prompt": prompt }) try: response = sagemaker_runtime.invoke_endpoint( EndpointName=endpoint_name, ContentType="application/json", Body=payload ) result = json.loads(response["Body"].read().decode()) return { "statusCode": 200, "headers": {"Content-Type": "application/json"}, "body": json.dumps(result) } except Exception as e: return { "statusCode": 500, "body": json.dumps({"error": str(e)}) }
2. 关键配置说明
- 权限配置:给Lambda执行角色添加以下IAM策略:
{ "Version": "2012-10-17", "Statement": [ { "Effect": "Allow", "Action": "sagemaker:InvokeEndpoint", "Resource": "arn:aws:sagemaker:你的区域:你的账号ID:endpoint/your-llava-endpoint-name" } ] }
- API Gateway设置:创建POST方法,集成Lambda,开启CORS(如果前端跨域调用),确保请求体可传递
image(base64字符串)和prompt参数。
内容的提问来源于stack exchange,提问作者Aleksandar Cvjetic
相关产品推荐
相关产品推荐

