如何用 Dask 集群运行 LightGBM 分布式训练

如何用 Dask 集群运行 LightGBM 分布式训练 如何用 Dask 集群运行 LightGBM 分布式训练【免费下载链接】LightGBMA fast, distributed, high performance gradient boosting (GBT, GBDT, GBRT, GBM or MART) framework based on decision tree algorithms, used for ranking, classification and many other machine learning tasks.项目地址: https://gitcode.com/GitHub_Trending/li/LightGBMLightGBM 的 Python 包内置了 Dask 集成lightgbm.dask自 3.2.0 版本引入macOS 支持自 4.7.0 版本引入可以用多台机器协作训练出一个完整的模型。这条路径适用于你手上已有或打算用dask.distributed搭建一个 Dask 集群、训练数据能切分成 Dask DataFrame / Dask Array / Dask Series 的场景。需要注意的适用边界是官方文档明确 Dask 集成只在 macOS 和 Linux 上测试过且 Python 必须为 64 位版本。安装带 Dask 依赖的 LightGBM在集群的每台节点上安装 LightGBM 的[dask]扩展它会一次性装齐lightgbm.dask所需的依赖pip install lightgbm[dask]示例代码还用到scikit-learn生成演示数据和dask.array前者可用pip install lightgbm[scikit-learn]补齐依赖安装细节见 python-package/README.rst。配置 Dask 集群线程与内存搭建集群时有两个官方明确给出的注意点。每个 worker 至少分配 2 个线程。如果线程不够通信任务和训练任务会互相阻塞训练可能明显变慢。如果没有其他进程与 Dask 竞争资源直接接受所选dask.distributed集群类型的默认nthreads即可from distributed import Client, LocalCluster cluster LocalCluster(n_workers3) client Client(cluster)训练期间关注 worker 内存。官方建议用 Dask 诊断 dashboard 或你常用的监控工具盯住 worker 内存占用。Dask worker 在内存压力过高时会自动把最近最少使用的数据溢写spill到磁盘磁盘 I/O 远慢于内存会显著拖慢计算。为降低触发内存上限的风险可以在执行任何数据加载或训练代码之前先重启一次 worker 进程client.restart()准备训练数据lightgbm.dask的估计器estimator要求矩阵或数组形态的数据以 Dask DataFrame、Dask Array 的形式提供部分场景也接受 Dask Series。建图过程是LightGBM 会把每个 worker 上持有的所有分区拼成一份本地数据集之后每个 Dask worker 上跑一个 LightGBM worker 进程参与分布式训练。切分数据时官方给出三条建议确保集群里每个 worker 都分到一部分训练数据尽量让每个 worker 持有的数据量大致相同数据集越小这一点越重要如果计划在同一份数据上训练多个模型例如调超参在训练前用client.persist()把数据物化一次避免重复计算。执行分布式训练以二分类为例最小可运行的主路径如下与 examples/python-guide/dask/binary-classification.py 一致示例中的数据量与树数量为文档示例值import dask.array as da from distributed import Client, LocalCluster from sklearn.datasets import make_blobs import lightgbm as lgb X, y make_blobs(n_samples1000, n_features50, centers2) cluster LocalCluster() client Client(cluster) dX da.from_array(X, chunks(100, 50)) dy da.from_array(y, chunks(100,)) dask_model lgb.DaskLGBMClassifier(n_estimators10) dask_model.fit(dX, dy) assert dask_model.fitted_fit()返回后检查dask_model.fitted_为True即表示本次分布式训练完成——这也是仓库示例脚本采用的验证方式。同样的模式适用于回归DaskLGBMRegressor见 examples/python-guide/dask/regression.py、多分类DaskLGBMClassifier[examples/python-guide/dask/multiclass-classification.py](https://link.gitcode.com/i/a4713d3c0e20507bcda1305e40077c5e)和排序DaskLGBMRanker[examples/python-guide/dask/ranking.py](https://link.gitcode.com/i/db362887006db149bd83e4d7b3dc742f)完整示例清单见 examples/python-guide/dask/README.md。Dask 估计器支持直接传入自定义目标函数4.0.0 起但要注意该函数会在每个 worker 进程上只对本地数据调用写函数时按此行为设计。集群网络受限制时指定端口训练开始时lightgbm.dask会在各 Dask worker 之间建立 LightGBM 网络worker 间通过 TCP socket 通信。默认使用随机可用端口如果集群里 worker 之间的通信被防火墙规则限制就必须显式告诉 LightGBM 用哪些端口。文档给出两种方式方式一machines参数逐 worker 指定地址:端口。逗号分隔每项对应一个 worker提供后 LightGBM 不再随机找端口import lightgbm as lgb machines 10.0.1.0:12401,10.0.2.0:12402,10.0.3.0:15000 dask_model lgb.DaskLGBMRegressor(machinesmachines)如果一台物理机上跑多个 Dask worker 进程例如nprocs2同一 IP 要给出多个不同端口的条目machines ,.join([ 10.0.1.0:16000, 10.0.1.0:16001, 10.0.2.0:16000, 10.0.2.0:16001, ]) dask_model lgb.DaskLGBMRegressor(machinesmachines)machines把网络细节完全交给你控制但代价是脆弱只要下列任一条件成立训练就会失败——machines中某个端口在训练开始时尚未开放训练数据分区所在机器没有出现在machines里machines里的某台机器不持有任何训练数据。方式二local_listen_port参数每台机器固定用同一个端口。前提是每台主机只跑一个 Dask worker 进程且你能确定一个在所有主机上都开放的端口。此时 LightGBM 只在该端口上工作并自动把网络限定在持有训练数据分区的 worker 集合内dask_model lgb.DaskLGBMRegressor(local_listen_port12400)这种方式比machines稍稳但如果local_listen_port在任一 worker 主机上未开放、或某台机器上同时跑了多个 Dask worker 进程训练同样会失败。local_listen_port的常规取值与含义见 docs/Parameters.rst。同一会话有多个 Dask client 时指定 client多数情况下不用指定默认使用distributed.default_client()返回的 client。只有当同一会话里有多个活跃 client例如在不同集群上跑多个训练任务时才通过估计器的client属性显式指定import lightgbm as lgb from distributed import Client, LocalCluster cluster LocalCluster() client Client(cluster) # 方式 1构造函数传参 dask_model lgb.DaskLGBMClassifier(clientclient) # 方式 2构造后再 set_params() dask_model lgb.DaskLGBMClassifier() dask_model.set_params(clientclient)训练完成后保存模型文档给出三个层次的保存选项按需选用直接 pickle Dask 估计器。用cloudpickle、joblib或pickle均可import pickle with open(dask-model.pkl, wb) as f: pickle.dump(dask_model, f) with open(dask-model.pkl, rb) as f: dask_model pickle.load(f)注意如果你之前显式设置了clientpickle 不会把它保存进去加载后如需指定 client用dask_model.set_params(clientclient)补上。转成 sklearn 估计器再存。训练完用to_local()得到等价的lightgbm.sklearn实例这样打分阶段就不依赖任何 Dask 库sklearn_model dask_model.to_local() # 类型为 lightgbm.sklearn.LGBMRegressor import joblib joblib.dump(sklearn_model, sklearn-model.joblib)提取底层 Booster。拿到dask_model.booster_之后可用序列化库保存也可以用bst.dump_model()导出为可写成 JSON 的字典、bst.model_to_string()导出为字符串、bst.save_model()写入文本文件。在 Dask 上做预测与评估.predict()接受 Dask Array 或 Dask DataFrame返回一个 Dask Array 的预测结果。示例见 examples/python-guide/dask/prediction.py训练后preds dask_model.predict(dX)得到分布式预测.compute()拉回客户端即可计算指标。该示例用sklearn.metrics计算 MSE但注释里明确这样做会把全部预测和目标值从 worker 拉回 client数据量大时应改用 dask-ml 的指标函数它们提供与sklearn.metrics等价的 API且指标计算全程分布式完成输入数据无需汇聚到单机。限制与不适用项Dask 集成只在 macOS 和 Linux 上测试过docs/Parallel-Learning-Guide.rst 与 python-package/README.rst 均给出此警告Windows 不在测试范围内。32 位 Python 不受支持安装前确认是 64 位版本。使用machines/local_listen_port的失败条件见上文属于文档明确列出的已知边界不是偶发问题。选择数据并行、特征并行还是投票并行tree_learnerdata/feature/voting官方建议按数据量与特征量的组合判断详见 docs/Parallel-Learning-Guide.rst 开头的选型表。完整参数与 API 细节可继续查阅 docs/Python-API.rst 中的 Dask API 条目和 docs/Parameters.rst。【免费下载链接】LightGBMA fast, distributed, high performance gradient boosting (GBT, GBDT, GBRT, GBM or MART) framework based on decision tree algorithms, used for ranking, classification and many other machine learning tasks.项目地址: https://gitcode.com/GitHub_Trending/li/LightGBM创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考