Scikit-learn训练任务:如何保存模型及数据集拆分并用于下游流水线?
在SageMaker中实现Scikit-learn模型+数据集拆分+指标的完整保存与下游复用
一、同时保存模型与数据集拆分
完全可以通过自定义训练脚本实现,SageMaker的Scikit-learn Estimator支持将模型以外的工件(比如拆分后的数据集)写入指定输出目录,具体操作如下:
- 利用SageMaker自动设置的环境变量:
SM_MODEL_DIR专门用于保存模型文件,SM_OUTPUT_DATA_DIR用于存储其他数据类工件(比如拆分后的训练/验证/测试集)。 - 训练脚本示例:
import os import pandas as pd import joblib from sklearn.model_selection import train_test_split from sklearn.ensemble import RandomForestClassifier # 加载训练数据 data = pd.read_csv(os.path.join(os.environ["SM_CHANNEL_TRAIN"], "train.csv")) X, y = data.drop("target", axis=1), data["target"] # 拆分数据集 X_train, X_temp, y_train, y_temp = train_test_split(X, y, test_size=0.3) X_val, X_test, y_val, y_test = train_test_split(X_temp, y_temp, test_size=0.5) # 训练模型 model = RandomForestClassifier(n_estimators=100) model.fit(X_train, y_train) # 保存模型到指定目录 joblib.dump(model, os.path.join(os.environ["SM_MODEL_DIR"], "model.joblib")) # 保存拆分后的数据集到输出目录 output_dir = os.environ["SM_OUTPUT_DATA_DIR"] dataset_files = [ (X_train, "train_data.csv"), (X_val, "val_data.csv"), (X_test, "test_data.csv"), (y_train, "train_labels.csv"), (y_val, "val_labels.csv"), (y_test, "test_labels.csv") ] for df, filename in dataset_files: df.to_csv(os.path.join(output_dir, filename), index=False)
- 启动Estimator时,通过
output_path指定S3存储路径,训练完成后,模型会存放在{output_path}/model下,数据集拆分文件则在{output_path}/output目录中。
二、工件用于下游流水线
这些工件完全可以接入下游训练或推理流水线:
- 模型工件:直接用于SageMaker推理端点、批量转换任务,或者在后续的模型调优、再训练步骤中加载使用;
- 数据集拆分文件:作为下游任务的输入源,比如模型验证流程、A/B测试、数据增强任务等,只需在下游步骤中指定对应的S3路径即可读取。
三、直接保存指标到JSON文件
不需要依赖日志正则匹配捕获指标,你可以在训练脚本中计算完指标后直接写入JSON文件到SM_OUTPUT_DATA_DIR:
import json from sklearn.metrics import accuracy_score, recall_score, f1_score # 计算验证集和测试集指标 val_metrics = { "accuracy": accuracy_score(y_val, model.predict(X_val)), "recall": recall_score(y_val, model.predict(X_val)), "f1": f1_score(y_val, model.predict(X_val)) } test_metrics = { "accuracy": accuracy_score(y_test, model.predict(X_test)), "recall": recall_score(y_test, model.predict(X_test)), "f1": f1_score(y_test, model.predict(X_test)) } # 保存到JSON文件 with open(os.path.join(output_dir, "validation_metrics.json"), "w") as f: json.dump(val_metrics, f, indent=2) with open(os.path.join(output_dir, "test_metrics.json"), "w") as f: json.dump(test_metrics, f, indent=2)
训练结束后,这些指标文件会和数据集拆分文件一起存储在S3,下游任务可以直接读取解析,用于可视化、报表生成或模型对比。
内容的提问来源于stack exchange,提问作者Rimon Dubnov
相关产品推荐
相关产品推荐

