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

SageMaker训练脚本无法加载预处理模型问题排查

问题根源

你误解了SageMaker中独立任务的目录机制:预处理和训练是两个独立的SKLearn作业,各自拥有专属的SM_MODEL_DIR(对应S3中的不同存储路径)。预处理任务把preprocessor.joblib存在了自己的模型输出路径,而训练任务默认只会加载自身配置的输入通道数据,不会自动拉取其他任务的输出文件。训练脚本里直接从自身的args.model_dir(即/opt/ml/model/)找预处理模型,自然找不到——这个路径里只有训练任务自己生成的文件,根本没有之前预处理任务保存的模型。

解决方案

方案1:给训练任务新增预处理模型输入通道(最简单的快速修复)

  1. 创建训练任务时,新增一个名为preprocessor的输入通道,指向预处理任务输出到S3的preprocessor.joblib所在路径。
  2. 修改训练脚本,从这个新输入通道加载模型:
    • 新增对应命令行参数,读取SageMaker的SM_CHANNEL_PREPROCESSOR环境变量
    • 调整加载逻辑,从preprocessor通道路径读取模型文件

方案2:合并预处理与训练为单个任务(适合小型实验)

如果不需要拆分任务,直接把两个脚本的逻辑合并成一个:先完成预处理拟合并保存到当前任务的SM_MODEL_DIR,接着直接加载这个模型进行训练,避免跨任务文件加载的问题。

方案3:用SageMaker Pipeline编排任务(生产级推荐)

通过SageMaker Pipeline定义任务依赖,让训练任务自动继承预处理任务的输出作为输入:

  • 定义ProcessingStep执行预处理,输出preprocessor工件
  • 定义TrainingStep,将预处理的输出指定为训练任务的输入通道
  • 这样训练任务就能直接从指定通道加载预处理模型,无需手动配置S3路径
修正后的训练脚本示例(方案1)
from sklearn.ensemble import RandomForestClassifier
from sklearn.metrics import accuracy_score, classification_report
import pandas as pd
import joblib
import os
import argparse

# 加载预处理模型函数
def load_preprocessor(model_dir):
    preprocessor_path = os.path.join(model_dir, "preprocessor.joblib")
    if os.path.exists(preprocessor_path):
        print(f"[INFO] Loading preprocessor from {preprocessor_path}")
        return joblib.load(preprocessor_path)
    else:
        raise FileNotFoundError(f"[ERROR] Preprocessor artifact not found at {preprocessor_path}")

if __name__ == "__main__":
    print("[INFO] Parsing arguments...")
    parser = argparse.ArgumentParser()
    parser.add_argument("--n_estimators", type=int, default=100)
    parser.add_argument("--random_state", type=int, default=42)
    parser.add_argument("--model-dir", type=str, default=os.environ.get("SM_MODEL_DIR"))
    parser.add_argument("--train", type=str, default=os.environ.get("SM_CHANNEL_TRAIN"))
    parser.add_argument("--test", type=str, default=os.environ.get("SM_CHANNEL_TEST"))
    # 新增预处理模型输入通道参数
    parser.add_argument("--preprocessor", type=str, default=os.environ.get("SM_CHANNEL_PREPROCESSOR"))
    parser.add_argument("--train-file", type=str, default="train.csv")
    parser.add_argument("--test-file", type=str, default="test.csv")
    args = parser.parse_args()

    print("[INFO] Loading preprocessor...")
    # 从预处理输入通道加载模型,而非训练任务自身的model_dir
    preprocessor = load_preprocessor(args.preprocessor)

    print("[INFO] Reading training and testing data...")
    train_df = pd.read_csv(os.path.join(args.train, args.train_file))
    test_df = pd.read_csv(os.path.join(args.test, args.test_file))

    # 定义特征与标签
    features = ['age', 'inactivity', 'recency', 'frequency']
    label = 'Risk'

    # 预处理数据
    X_train = preprocessor.transform(train_df[features])
    y_train = train_df[label]
    X_test = preprocessor.transform(test_df[features])
    y_test = test_df[label]

    print("[INFO] Training Random Forest model...")
    model = RandomForestClassifier(n_estimators=args.n_estimators, random_state=args.random_state)
    model.fit(X_train, y_train)

    # 保存训练好的模型
    model_path = os.path.join(args.model_dir, "random_forest_model.joblib")
    joblib.dump(model, model_path)
    print(f"[INFO] Trained model saved as random_forest_model.joblib at {model_path}")

    print("[INFO] Evaluating model on test data...")
    y_pred = model.predict(X_test)
    acc = accuracy_score(y_test, y_pred)
    report = classification_report(y_test, y_pred)

    print("Accuracy:", acc)
    print("Classification Report:\n", report)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 22:04:53