SageMaker训练脚本无法加载预处理模型问题排查
问题根源
你误解了SageMaker中独立任务的目录机制:预处理和训练是两个独立的SKLearn作业,各自拥有专属的SM_MODEL_DIR(对应S3中的不同存储路径)。预处理任务把preprocessor.joblib存在了自己的模型输出路径,而训练任务默认只会加载自身配置的输入通道数据,不会自动拉取其他任务的输出文件。训练脚本里直接从自身的args.model_dir(即/opt/ml/model/)找预处理模型,自然找不到——这个路径里只有训练任务自己生成的文件,根本没有之前预处理任务保存的模型。
解决方案
方案1:给训练任务新增预处理模型输入通道(最简单的快速修复)
- 创建训练任务时,新增一个名为
preprocessor的输入通道,指向预处理任务输出到S3的preprocessor.joblib所在路径。 - 修改训练脚本,从这个新输入通道加载模型:
- 新增对应命令行参数,读取SageMaker的
SM_CHANNEL_PREPROCESSOR环境变量 - 调整加载逻辑,从
preprocessor通道路径读取模型文件
- 新增对应命令行参数,读取SageMaker的
方案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
相关产品推荐
相关产品推荐

