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

SageMaker中lightgbm-classification-model脚本模式无法导入LightGBM及脚本要求咨询

问题解决与训练脚本规范说明

一、解决ModuleNotFoundError: No module named 'lightgbm'

针对你使用的lightgbm-classification-model:2.1.3容器,可通过两种方式安装依赖:

  • 在src目录下创建requirements.txt,写入lightgbm==2.1.3,然后在定义Estimator时指定source_dir='src',SageMaker会自动安装该文件内的依赖包。
  • 直接在Estimator的pip_packages参数中添加["lightgbm==2.1.3"],训练启动时会自动安装指定版本的lightgbm。

二、SageMaker训练入口脚本(train.py)规范

结合LightGBM分类场景,脚本需遵循以下通用规则:

1. 预期输入

  • 数据路径:SageMaker会将训练/验证数据挂载到容器固定路径,默认:
    • 训练数据:/opt/ml/input/data/train(对应Estimator中channel='train'的TrainingInput)
    • 验证数据:/opt/ml/input/data/validation(若指定验证通道)
  • 超参数:可通过命令行参数解析,或读取os.environ['SM_HP_<超参数名>']环境变量获取(比如定义的hyperparameters={"num_round":100},可通过SM_HP_NUM_ROUND读取)。

2. 预期输出

训练好的模型必须保存到/opt/ml/model目录下,SageMaker会自动将该目录内容打包上传至指定S3路径:

  • LightGBM模型通常保存为model.txt或.bst格式,直接放入/opt/ml/model即可。
  • 日志、评估指标等非模型文件可写入/opt/ml/output,但该目录内容不会被自动保存,需自行上传至S3。

3. 脚本结构与函数规范

脚本无需强制特定函数签名,但通常按以下流程编写:

import os
import argparse
import lightgbm as lgb
import pandas as pd

def main():
    # 解析参数:包含SageMaker自动传入的路径参数与自定义超参数
    parser = argparse.ArgumentParser()
    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('--validation', type=str, default=os.environ.get('SM_CHANNEL_VALIDATION'))
    parser.add_argument('--num_round', type=int, default=100)
    parser.add_argument('--learning_rate', type=float, default=0.1)
    args = parser.parse_args()

    # 加载并预处理训练数据
    train_df = pd.read_csv(os.path.join(args.train, 'train.csv'))
    train_set = lgb.Dataset(train_df.drop('label', axis=1), label=train_df['label'])

    # 加载验证数据(可选)
    val_set = None
    if args.validation:
        val_df = pd.read_csv(os.path.join(args.validation, 'val.csv'))
        val_set = lgb.Dataset(val_df.drop('label', axis=1), label=val_df['label'], reference=train_set)

    # 定义模型参数并训练
    params = {
        'objective': 'binary',  # 按需改为multiclass等分类类型
        'metric': 'auc',
        'learning_rate': args.learning_rate,
        'num_leaves': 31
    }
    model = lgb.train(
        params,
        train_set,
        num_boost_round=args.num_round,
        valid_sets=[val_set] if val_set else None,
        early_stopping_rounds=10
    )

    # 保存模型到指定路径
    model.save_model(os.path.join(args.model_dir, 'model.txt'))

if __name__ == '__main__':
    main()

4. 常用环境变量

SageMaker会自动设置以下环境变量,可直接读取:

  • SM_MODEL_DIR:模型必须保存的路径
  • SM_CHANNEL_TRAIN/SM_CHANNEL_VALIDATION:训练/验证数据的挂载路径
  • SM_HP_<超参数名>:自定义超参数的环境变量形式
  • SM_NUM_GPUS:实例可用的GPU数量(若使用GPU实例)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 15:00:01