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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 03:47:31