You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用SageMaker SDK部署自定义SKlearn Pipeline模型遇依赖加载错误

解决SageMaker部署自定义Transformer类的依赖问题

核心问题原因

你遇到的报错是因为joblib加载模型时找不到RecodeCategorias类的定义——训练时如果类是在主脚本(__main__模块)里定义的,保存后的模型会把类关联到__main__,但SageMaker推理环境的__main__是gunicorn进程,自然找不到这个类。

正确的依赖传递与部署步骤

1. 规范目录结构

把自定义类单独放到一个模块文件里,和推理脚本、模型文件同目录:

model_package/
├── inference.py          # 推理逻辑脚本
├── custom_transformers.py # 存放RecodeCategorias类的定义
├── model.joblib
└── pipeline.joblib

2. 确保训练时类的导入路径正确

训练阶段,不要在主训练脚本里直接定义RecodeCategorias,而是从custom_transformers.py导入:

# 训练脚本示例
from custom_transformers import RecodeCategorias
from sklearn.pipeline import Pipeline
from sklearn.linear_model import LogisticRegression

# 构建Pipeline
pipeline = Pipeline([
    ('recode', RecodeCategorias()),
    ('clf', LogisticRegression())
])

# 训练并保存
pipeline.fit(X_train, y_train)
joblib.dump(pipeline, 'pipeline.joblib')

3. 推理脚本正确导入自定义类

在inference.py里明确从自定义模块导入类:

# inference.py
from custom_transformers import RecodeCategorias
import joblib
import os

def model_fn(model_dir):
    # 加载整个Pipeline
    pipeline_path = os.path.join(model_dir, 'pipeline.joblib')
    return joblib.load(pipeline_path)

# 可选:如果需要自定义input_fn/output_fn,按需添加

4. 打包并上传模型

将整个model_package目录打包成model.tar.gz:

cd model_package
tar -czf model.tar.gz *

然后把这个压缩包上传到你的S3桶。

5. SageMaker部署配置

使用SKLearnModel部署时,直接指定model_data为S3上的model.tar.gz路径即可,无需额外配置dependencies或复杂的source_dir(因为所有依赖文件都在压缩包内):

from sagemaker.sklearn.model import SKLearnModel

model = SKLearnModel(
    model_data='s3://your-bucket/path/to/model.tar.gz',
    role='your-sagemaker-role',
    framework_version='1.2-1'  # 匹配你训练时的SKlearn版本
)

predictor = model.deploy(
    initial_instance_count=1,
    instance_type='ml.t2.medium'
)

关键注意点

  • 训练和推理环境的SKlearn版本必须一致,避免兼容性问题。
  • 不要把自定义类放在子目录里(除非你在inference.py里添加相对路径导入,比如from .custom_transformers import ...,但同目录更简单)。
  • 打包时确保所有必要文件都被包含,不要遗漏custom_transformers.py。

内容的提问来源于stack exchange,提问作者Geo

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.27 17:42:06