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
相关产品推荐
相关产品推荐

