在Amazon SageMaker中部署TF2模型时如何编写自定义推理脚本
SageMaker部署TF2模型时加载额外工具的解决方案
问题根源
SageMaker的TensorFlowModel默认采用TensorFlow Serving进行推理,不会调用自定义脚本中的model_fn——这是它和HuggingFace/Sklearn模型的核心区别。要加载输入映射JSON、Transformer分词器这类额外工具,必须切换到Python自定义推理模式,让SageMaker执行你定义的完整推理逻辑。
修改后的自定义推理脚本
调整脚本,在model_fn中同时加载模型和所有额外工具,并在输入处理环节调用这些工具:
import tensorflow as tf import numpy as np import json from transformers import AutoTokenizer # 全局变量存储模型、输入映射、分词器 model = None input_mapping = None tokenizer = None def model_fn(model_dir): """加载模型及所有额外依赖工具""" global model, input_mapping, tokenizer # 加载TF2 SavedModel model = tf.keras.models.load_model(f"{model_dir}/saved_model") # 加载输入映射JSON文件 with open(f"{model_dir}/input_mapping.json", "r") as f: input_mapping = json.load(f) # 加载预训练Transformer分词器(假设分词器文件已打包到model_dir下的tokenizer目录) tokenizer = AutoTokenizer.from_pretrained(f"{model_dir}/tokenizer") return model def input_fn(input_data, content_type): """结合分词器和输入映射处理请求输入""" if content_type == "application/json": input_dict = json.loads(input_data) raw_texts = input_dict["instances"] # 使用分词器对文本进行编码 encoded_inputs = tokenizer( raw_texts, padding=True, truncation=True, return_tensors="tf" ) # 根据输入映射转换为模型需要的输入格式 model_input = {input_mapping[k]: v for k, v in encoded_inputs.items()} return model_input else: raise ValueError(f"不支持的内容类型: {content_type}") def predict_fn(input_data, model): """执行模型推理""" return model.predict(input_data) def output_fn(prediction, accept): """序列化推理结果""" if accept == "application/json": response = {"predictions": prediction.tolist()} return json.dumps(response), "application/json" else: raise ValueError(f"不支持的响应类型: {accept}") # 注意:Python自定义推理模式下,无需input_handler和output_handler,SageMaker会直接调用上述四个核心函数
部署代码调整
部署时需明确指定自定义脚本为入口点,并确保额外文件与模型一起打包上传:
from sagemaker.tensorflow import TensorFlowModel # 初始化TF模型,指定入口脚本和依赖 model = TensorFlowModel( model_data="s3://your-bucket/model.tar.gz", # 包含模型、脚本、额外文件的压缩包 role="your-sagemaker-iam-role", framework_version="2.15", # 匹配你的TF版本 entry_point="inference.py", # 自定义推理脚本的文件名 dependencies=["requirements.txt"] # 额外依赖包配置文件 ) # 部署推理端点 predictor = model.deploy( initial_instance_count=1, instance_type="ml.t2.medium" )
关键注意事项
- 打包要求:将
inference.py、requirements.txt、input_mapping.json、分词器文件夹、TF SavedModel目录一起压缩成model.tar.gz,上传到S3。 - 依赖配置:
requirements.txt需包含所有依赖包,例如:transformers==4.35.2 tensorflow==2.15.0 numpy==1.24.3 - 目录结构:压缩包内的目录结构需清晰,示例:
model.tar.gz/ ├── saved_model/ # TF SavedModel的标准目录 ├── inference.py ├── requirements.txt ├── input_mapping.json └── tokenizer/ # 包含分词器的vocab.json、tokenizer_config.json等文件
内容的提问来源于stack exchange,提问作者shreyansh
相关产品推荐
相关产品推荐

