如何在本地加载SageMaker导出至S3的模型并运行预测?
解决方案:加载SageMaker导出的.tar.gz模型到本地内存运行预测
我完全懂你的困扰——SageMaker导出的模型包确实和它的容器生态绑定得很紧,刚拿到手直接解压,要么找不到熟悉的模型文件,要么加载了也跑不起来。别慌,我给你一步步拆解可行的方案,核心思路是先搞清楚模型包的内部结构,再用对应框架的原生方法加载。
第一步:把S3上的模型包下载到内存/本地
首先用boto3把S3上的.tar.gz包下载到内存对象里(不用存到本地磁盘也可以):
import boto3 import io # 初始化S3客户端 s3 = boto3.client('s3') # 替换成你的Bucket名称和模型文件路径 bucket_name = "your-s3-bucket-name" model_s3_key = "path/to/your/model.tar.gz" # 下载到内存BytesIO对象 response = s3.get_object(Bucket=bucket_name, Key=model_s3_key) model_tar_bytes = io.BytesIO(response['Body'].read())
第二步:查看模型包的内部结构
SageMaker导出的模型包结构取决于你训练时用的框架(PyTorch/TensorFlow/XGBoost等),先解压看看里面有什么:
import tarfile import tempfile import os with tempfile.TemporaryDirectory() as tmp_dir: # 解压内存中的tar包到临时目录 with tarfile.open(fileobj=model_tar_bytes, mode='r:gz') as tar: tar.extractall(path=tmp_dir) # 列出目录内容,找到核心模型文件 print("模型包内文件结构:") for root, dirs, files in os.walk(tmp_dir): level = root.replace(tmp_dir, '').count(os.sep) indent = ' ' * 2 * level print(f"{indent}{os.path.basename(root)}/") sub_indent = ' ' * 2 * (level + 1) for file in files: print(f"{sub_indent}{file}")
运行这段代码后,你就能看到模型包的具体结构了——比如PyTorch通常是model.pth/model.pt,TensorFlow是带saved_model.pb的目录,XGBoost是xgboost-model文件。
第三步:根据框架类型加载模型并预测
下面针对主流框架给出具体的加载代码:
情况1:PyTorch模型
如果看到model.pth或model.pt文件,直接用PyTorch原生方法加载:
import torch # 假设你已经定义了训练时的模型结构(必须和训练时一致!) class YourModel(torch.nn.Module): def __init__(self): super().__init__() # 替换成你的模型层定义 self.fc = torch.nn.Linear(10, 2) def forward(self, x): return self.fc(x) with tempfile.TemporaryDirectory() as tmp_dir: with tarfile.open(fileobj=model_tar_bytes, mode='r:gz') as tar: tar.extractall(path=tmp_dir) # 替换成你的模型文件路径(从第二步的输出里找) model_path = os.path.join(tmp_dir, "model.pth") # 加载模型并切换到评估模式 model = YourModel() model.load_state_dict(torch.load(model_path)) model.eval() # 运行预测(替换成你的输入数据) input_tensor = torch.randn(1, 10) # 示例输入 with torch.no_grad(): output = model(input_tensor) print("预测结果:", output)
情况2:TensorFlow/Keras模型
如果看到包含saved_model.pb和variables/目录的SavedModel结构,用TensorFlow原生方法加载:
import tensorflow as tf with tempfile.TemporaryDirectory() as tmp_dir: with tarfile.open(fileobj=model_tar_bytes, mode='r:gz') as tar: tar.extractall(path=tmp_dir) # 自动找到SavedModel目录(找包含saved_model.pb的文件夹) saved_model_dir = None for root, _, files in os.walk(tmp_dir): if "saved_model.pb" in files: saved_model_dir = root break # 加载模型 model = tf.saved_model.load(saved_model_dir) # 获取默认的推理签名(SageMaker默认用serving_default) infer_fn = model.signatures["serving_default"] # 运行预测(替换成你的输入数据) input_tensor = tf.random.normal([1, 28, 28]) # 示例输入 output = infer_fn(input_tensor) print("预测结果:", output)
情况3:XGBoost模型
如果看到xgboost-model文件,用XGBoost原生方法加载:
import xgboost as xgb with tempfile.TemporaryDirectory() as tmp_dir: with tarfile.open(fileobj=model_tar_bytes, mode='r:gz') as tar: tar.extractall(path=tmp_dir) # 替换成你的模型文件路径 model_path = os.path.join(tmp_dir, "xgboost-model") # 加载模型 model = xgb.Booster() model.load_model(model_path) # 运行预测(输入需要是DMatrix格式) input_data = xgb.DMatrix([[0.1, 0.2, 0.3, 0.4]]) # 示例输入 output = model.predict(input_data) print("预测结果:", output)
关键注意事项
- 模型结构必须一致:加载PyTorch模型时,必须提前定义好和训练时完全一样的模型类,否则会加载失败。
- 预处理/后处理自己实现:SageMaker容器里的推理脚本通常包含数据预处理和后处理逻辑,脱离容器后你需要自己在代码中实现这些步骤。
- 依赖版本匹配:本地环境的框架版本(比如torch/tf/xgboost的版本)要和训练时SageMaker容器的版本尽量一致,避免兼容性问题。
内容的提问来源于stack exchange,提问作者NotSoShabby
相关产品推荐
相关产品推荐

