AWS无服务器SageMaker端点调用HuggingFace摘要模型报错求助
问题分析与解决方案
错误根源
你遇到的TypeError: 'str' object is not callable错误,本质是SageMaker容器加载模型时,transformers pipeline的tokenizer属性被赋值为字符串而非可调用的PreTrainedTokenizer实例。
原因在于你使用的**新版本推理镜像(1.13.1-transformers4.26.0)**与AWS JumpStart提供的旧版摘要模型包(infer-huggingface-summarization-bert-small2bert-small-finetuned-cnn-daily-mail-summarization.tar.gz)的加载逻辑不兼容。对比可正常运行的翻译模型,它使用的是旧版镜像(1.7.1-transformers4.6.1),该版本的加载逻辑与旧模型包完全匹配。
修复方案
方案1:更换兼容的镜像版本
直接将摘要模型的镜像替换为与翻译模型同版本的旧镜像,确保模型包加载逻辑兼容:
SageMakerModel: Type: AWS::SageMaker::Model Properties: ModelName: SummarizationModel Containers: - Image: "763104351884.dkr.ecr.us-east-1.amazonaws.com/huggingface-pytorch-inference:1.7.1-transformers4.6.1-cpu-py36-ubuntu18.04" ModelDataUrl: "s3://jumpstart-cache-prod-us-east-1/huggingface-infer/infer-huggingface-summarization-bert-small2bert-small-finetuned-cnn-daily-mail-summarization.tar.gz" Mode: SingleModel ExecutionRoleArn: !GetAtt SageMakerExecutionRole.Arn
方案2:自定义推理脚本(保留新镜像)
如果必须使用新版镜像,需要自定义模型加载脚本,手动初始化模型和tokenizer:
- 下载原模型包到本地:
aws s3 cp s3://jumpstart-cache-prod-us-east-1/huggingface-infer/infer-huggingface-summarization-bert-small2bert-small-finetuned-cnn-daily-mail-summarization.tar.gz ./model.tar.gz tar -xvf model.tar.gz
- 创建自定义
inference.py:
from transformers import pipeline, AutoModelForSeq2SeqLM, AutoTokenizer def model_fn(model_dir): tokenizer = AutoTokenizer.from_pretrained(model_dir) model = AutoModelForSeq2SeqLM.from_pretrained(model_dir) return pipeline("summarization", model=model, tokenizer=tokenizer) def predict_fn(input_data, model): return model(input_data["inputs"], max_length=130, min_length=30, do_sample=False)
- 重新打包模型:
tar -czvf custom-model.tar.gz inference.py config.json pytorch_model.bin tokenizer.json tokenizer_config.json vocab.txt
- 将打包后的模型上传到你的S3桶,更新CloudFormation中的
ModelDataUrl为新的S3路径。
正确Payload格式与调用示例
确认模型加载正常后,使用标准格式的Payload即可(无需添加"Summarize this text:"前缀):
import boto3 import json client = boto3.client('runtime.sagemaker') payload = { "inputs": "This is a beautiful day. I am happy. I am going to the park." } response = client.invoke_endpoint( EndpointName="SummarizationEndpoint", ContentType="application/json", Accept="application/json", Body=json.dumps(payload) ) # 解析响应 result = json.loads(response['Body'].read().decode()) print(result)
内容的提问来源于stack exchange,提问作者Stefan Radulian
相关产品推荐
相关产品推荐

