尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

Angel 上的分布式 GBDT(梯度提升决策树):参数服务器实现、两阶段分裂算法与实战配置

Angel 上的分布式 GBDT(梯度提升决策树):参数服务器实现、两阶段分裂算法与实战配置 人工智能机器学习分布式训练图计算后端【免费下载链接】angelA Flexible and Powerful Parameter Server for large-scale machine learning项目地址https://gitcode.com/gh_mirrors/an/angel点击查看免费下载本文以 Angel 开源仓库的官方算法文档 gbdt_on_angel.md 为核心骨架结合 GBDT 源码实现 与单元测试系统讲解 GBDTGradient Boosting Decision Tree梯度提升决策树在 Angel 参数服务器架构下的分布式训练原理、梯度直方图存储方案、两阶段分裂算法以及完整的训练/预测提交命令与调参指南。读完本文你将掌握如何在 Angel 上提交并调优一个分布式 GBDT 训练任务并理解其相对 Spark 版 GBDT、MPI 版 XGBoost 的性能优势来源。1. GBDT 算法核心思想GBDT 是一种集成学习算法通过串行地训练多棵决策树弱分类器将每棵树的预测结果累加起来从而不断提升整体模型的分类或回归精度。它在大量分类和回归场景中都有不错的效果。如下图所示这是一个对一群消费者的消费力进行预测的例子其处理流程为第一棵树根节点分裂根节点选取的特征是「年龄」年龄小于 30 的被分到左子节点年龄大于 30 的被分到右叶子节点右叶子节点的预测值为 1第一棵树左子节点继续分裂左一节点继续分裂分裂特征是「月薪」月薪小于 10K 划为左叶子节点预测值为 5月薪大于 10K 划为右叶子节点预测值为 10更新预测值建立完第一棵树之后C、D 和 E 的预测值被更新为 1A 为 5B 为 10建立第二棵树根据新的预测值残差开始建立第二棵树第二棵树的根节点分裂特征是「性别」女性预测值为 0.5男性预测值为 1.5累加预测值建立完第二棵树之后将第二棵树的预测值加到每个消费者已有的预测值上。例如 A 的预测值为两棵树预测值之和5 0.5 5.5迭代优化通过这种方式逐棵树地拟合残差不断地优化预测准确率。从源码看这种逐棵建树、累加预测的过程由 GBDTController.java 中的状态机驱动每个 Worker 维护一棵RegTree森林RegTree[] forest建树完成后通过updateInsPreds()将叶子权重乘以学习率累加到每条样本的预测值上对应源码trainDataStore.preds[insIdx] this.param.learningRate * weight与上述算法流程完全一致。2. GBDT 在 Angel 上的分布式实现2.1 参数服务器中的参数存储为了优化算法性能Angel 将 GBDT 训练过程中反复更新和传递的全部参数以矩阵形式存储在参数服务器PS上主要包括每个树节点的分裂特征 IDfeature ID每个树节点的分裂特征值feature Value叶子节点的预测值leaf-prediction全局一阶梯度直方图grad histogram全局二阶梯度直方图hess histogram这些参数矩阵在整个 GBDT 的计算过程中会被反复更新和传递。在仓库源码 GBDTModel.scala 中这些参数被落实为 10 个 PS 矩阵每个矩阵均按 PS 节点数做了列切分以解决高维模型汇总的单点瓶颈矩阵名源码常量维度行 × 列行类型用途Quantile sketchgbdt.sketch1 × 特征数 × 分裂数T_DOUBLE_DENSE存储每个特征的候选分裂值候选分裂点Sampled featuregbdt.feature.sample树数 × 采样特征数T_INT_DENSE记录每棵树采样出的特征子集Grad/Heiss histogramgbdt.grad.histogram.node{i}1 × 2×分裂数×采样特征数T_DOUBLE_DENSE每个树节点一张存储全局一阶/二阶梯度直方图Active tree nodesgbdt.active.nodes1 × 最大节点数T_INT_DENSE标记哪些树节点处于待分裂活跃状态Split featuregbdt.split.feature树数 × 最大节点数T_INT_DENSE每个树节点的分裂特征 IDSplit valuegbdt.split.value树数 × 最大节点数T_DOUBLE_DENSE每个树节点的分裂特征值Split gaingbdt.split.gain树数 × 最大节点数T_DOUBLE_DENSE每个树节点分裂的目标函数增益Node grad statsgbdt.node.grad.stats树数 × 2×最大节点数T_DOUBLE_DENSE每个节点的梯度统计一阶和、二阶和Node predictgbdt.node.predict树数 × 最大节点数T_DOUBLE_DENSE叶子节点的预测值Categorical featuregbdt.feature.categoryWorker 数 × 类别特征数×分裂数T_DOUBLE_DENSE类别特征的分裂点其中分裂直方图、梯度直方图、活跃节点等中间矩阵均通过.setNeedSave(false)标记为不落盘仅保留分裂特征、分裂值、节点预测值等最终模型参数。模型定义中还根据indexRange % psNumber对特征数做了向上取整调整保证参数矩阵能被均匀切分到各个 PS 节点上。2.2 整体训练流程GBDT 的训练包含几大步骤计算候选分裂点Create sketch扫描训练数据对每种特征计算候选分裂特征值从而得到候选分裂点分裂特征 分裂特征值常用的方法有 Quantile sketch。Angel 实现中由 TYahooSketchSplit.java 完成并由领导 WorkertaskIndex 0将全局 sketch 推送到 PS 的gbdt.sketch矩阵上创建决策树New treeWorker 创建新的树进行初始化工作包括初始化树结构、计算训练数据的一阶和二阶梯度、初始化一个待处理树节点的队列、将树的根节点加入队列。对应源码createNewTree()新建RegTree、重置活跃节点、把根节点nid0设为活跃、调用calGradPairs()计算梯度寻找最佳分裂点 分裂树节点这是最关键的一步下文单独展开计算合并叶子节点的预测值Finish treeWorker 计算出叶子节点的预测值并推送给 PSupdateLeafPreds()写入gbdt.node.predict完成一颗决策树重新开始第 2 步直到训练完所有决策树计算并输出性能指标准确率、误差等输出训练模型。在 GBDTController.java 中上述步骤被建模为一个状态机CREATE_SKETCH → GET_SKETCH → SAMPLE_FEATURE → NEW_TREE → RUN_ACTIVE → FIND_SPLIT → AFTER_SPLIT → FINISH_TREE → FINISHED每一轮通过updatePhase()推进直到currentTree treeNum完成全部树的训练。2.3 最佳分裂点搜索与两阶段分裂GBDT 的精髓如何寻找最佳分裂点并进行分裂是 GBDT 的精髓和难点也是参数服务器对它产生重要价值的所在。整体流程为计算梯度直方图Run active node从待处理树节点的队列中取出待处理的树节点在 Worker 上根据本节点的训练数据计算局部梯度直方图包括一阶和二阶。源码中通过HistCalThread按batchSize分批并行计算并支持HistSubThread用父节点直方图 - 兄弟节点直方图的直方图减法来快速得到另一子节点的直方图ml.gbdt.hist.subtraction默认开启大幅减少重复计算同步 合并直方图Worker 通过 PS 接口将局部梯度直方图推送到参数服务器。在发送之前每个局部梯度直方图被切分为 P 个分块P 为参数服务器节点个数每个分块分别发送到对应的参数服务器节点PS 节点接收到 Worker 发送的局部梯度直方图后确定处理的树节点将其累加到对应的全局梯度直方图上。源码中每个节点直方图独立存放在gbdt.grad.histogram.node{nid}矩阵中按2 * splitNum * sampleFeatNum / psNumber做列切分寻找最佳分裂点Find splitWorker 使用参数服务器提供的计算最佳分裂点的接口从参数服务器获取每个参数服务器节点上的最佳分裂点然后比较 P 个分裂点的目标函数增益选取增益最大的分裂点作为全局最佳分裂点。源码中多个活跃树节点在各 Task 间采用 Round-Robin 分配taskContext.getTaskIndex() activeTNodeNum保证分裂计算并行均衡分裂树节点After splitWorker 根据计算得到的最佳分裂点创建叶子节点将本节点的训练数据切分到两个叶子节点上updateTrainInsPos()使用快速划分原地重排样本区间如果树的高度没有达到最大限制则将两个叶子节点加入到待处理树节点的队列。其中两阶段分裂算法two-phase tree splittingml.gbdt.server.split是关键优化当开关开启时Worker 不再把整棵直方图拉回本地计算而是通过 PS 上的 PSFParameter Server Function参数服务器函数GBDTGradHistGetRowFunc直接在 PS 端并行扫描直方图、找到每个分块内的局部最佳分裂点再把 P 个局部最佳分裂点而不是庞大的直方图返回给 WorkerWorker 只需比较增益即可选出全局最佳分裂点。对应源码位于 GBDTGradHistGetRowFunc.java 与findSplit()中isServerSplit分支。从上面的算法逻辑剖析可以看出GBDT 算法存在大量的模型更新和同步操作非常适合参数服务器的系统架构具体体现在超大模型GBDT 用到的梯度直方图的大小与特征数量成正比对于高维大数据集梯度直方图会非常大。Angel 将梯度直方图切分到多个 PS 节点上存储有效解决了高维度模型在汇总参数时的单点瓶颈问题两阶段树分裂算法在寻找最佳分裂点时在多个 PS 节点上并行处理只需要将局部最佳分裂点返回给 Worker通信开销几乎可以忽略不计。整体来看Angel 的 PS 优势使得它在该算法上的性能远超 Spark 版本的 GBDT 实现也显著优于 MPI 版本的 XGBoost。3. 运行 GBDT输入格式与参数详解3.1 输入格式ml.feature.index.range特征向量的维度ml.data.type支持dummy、libsvm两种数据格式具体参考 Angel 数据格式说明。仓库内置的 GBDT 单元测试 GBDTTest.java 使用agaricus_127d_train.libsvm127 维特征作为训练数据、agaricus_127d_test.libsvm作为预测数据该数据已随仓库提供在 data/agaricus 目录下可直接用于本地复现。3.2 算法参数参数含义默认值源码 MLConf说明ml.gbdt.tree.num树的数量10默认 10增大可提升精度但增加训练时间ml.gbdt.tree.depth树的最大高度5默认 5深度越大模型越复杂ml.gbdt.split.num每个特征的分裂点的数量10决定候选分裂点个数也决定直方图长度ml.learn.rate学习速率0.5每棵树预测值的缩放系数一般取 0.010.2ml.data.validate.ratio每次 validation 的样本比率设为 0 时不做 validation0.05训练过程中按比例切分验证集ml.gbdt.sample.ratio特征下采样的比率1默认 1 表示使用全部特征ml.gbdt.server.split两阶段分裂算法开关false开启后在 PS 端并行寻找最佳分裂点ml.gbdt.batch.size并行训练时一个批量的数量10000控制直方图计算时分批粒度angel.compress.bytes低精度压缩每个浮点数的大小8可设为 [1,8]如设为 2 即用 2 字节量化梯度再传输补充说明来自源码 MLConf.scala 与 GBDTParam.javaml.gbdt.task.type可设为classification默认使用二元逻辑损失binary:logistic或regression使用平方损失回归任务对应的评估指标为 RMSEml.gbdt.cate.feat用于声明类别特征另外还有ml.gbdt.reg.alpha/ml.gbdt.reg.lambda分裂增益的正则项、ml.gbdt.min.child.weight子节点最小二阶梯度权重和、ml.gbdt.max.node.num最大节点数、ml.gbdt.feature.sample.ratio等扩展参数可供精细调优。当angel.compress.bytes 1 || 8时源码会判定为非法配置并回退到默认值 8。3.3 输入输出参数参数含义angel.train.data.path训练数据的输入路径angel.predict.data.path预测数据的输入路径ml.gbdt.cate.feat类别特征特征id:特征范围的格式以逗号分隔例如0:2,1:3。设为none表示没有离散特征设为all表示全部为离散特征ml.model.type模型类型默认为T_FLOAT_DENSEangel.save.model.path训练完成后模型的保存路径angel.predict.out.path预测结果的保存路径angel.log.path日志文件的保存路径3.4 资源参数参数含义默认值angel.workergroup.numberWorker 个数-angel.worker.memory.gbWorker 申请内存大小单位 GB-angel.worker.task.number每个 Worker 上的 task 的个数1angel.ps.numberPS 个数1angel.ps.memory.gbPS 申请内存大小单位 GB-3.5 训练任务启动命令示例angel-submit \ -Dangel.am.log.levelINFO \ -Dangel.ps.log.levelINFO \ -Dangel.worker.log.levelINFO \ -Dangel.app.submit.classcom.tencent.angel.ml.GBDT.GBDTRunner \ -Daction.typetrain \ -Dml.data.typelibsvm \ -Dml.model.typeT_FLOAT_DENSE \ -Dml.data.validate.ratio0.1 \ -Dml.feature.index.range10000 \ -Dml.gbdt.cate.featnone \ -Dml.gbdt.tree.num20 \ -Dml.gbdt.tree.depth7 \ -Dml.gbdt.split.num10 \ -Dml.gbdt.sample.ratio1.0 \ -Dml.learn.rate0.01 \ -Dml.gbdt.server.splittrue \ -Dangel.compress.bytes2 \ -Dangel.train.data.path$input_path \ -Dangel.save.model.path$model_path \ -Dangel.workergroup.number50 \ -Dangel.worker.memory.gb10 \ -Dangel.task.data.storage.levelmemory \ -Dangel.worker.task.number1 \ -Dangel.ps.number50 \ -Dangel.ps.memory.gb103.6 预测任务启动命令示例angel-submit \ -Dangel.am.log.levelINFO \ -Dangel.ps.log.levelINFO \ -Dangel.worker.log.levelINFO \ -Dangel.app.submit.classcom.tencent.angel.ml.GBDT.GBDTRunner \ -Daction.typepredict \ -Dml.data.typelibsvm \ -Dml.model.typeT_FLOAT_DENSE \ -Dml.data.validate.ratio0.1 \ -Dml.feature.index.range10000 \ -Dml.gbdt.tree.num20 \ -Dml.gbdt.tree.depth7 \ -Dml.gbdt.sample.ratio1.0 \ -Dml.learn.rate0.01 \ -Dangel.predict.data.path$input_path \ -Dangel.save.model.path$model_path \ -Dangel.predict.out.path$predict_path \ -Dangel.workergroup.number50 \ -Dangel.worker.memory.gb10 \ -Dangel.task.data.storage.levelmemory \ -Dangel.worker.task.number1 \ -Dangel.ps.number50 \ -Dangel.ps.memory.gb103.7 任务执行链路源码视角训练与预测任务的入口为 GBDTRunner.scala。训练模式下Runner 依次执行创建 AngelClient →startPSServer()启动参数服务器 →loadModel()加载或新建GBDT 模型 →runTask(classOf[GBDTTrainTask])在 Worker 上拉起训练 Task →waitForCompletion()等待训练完成 →saveModel()保存模型 →stop()释放集群资源。预测模式则额外将矩阵传输超时angel.worker.matrix.transfer.request.timeout.ms调大至 60000ms 后执行runTask(classOf[GBDTPredictTask])。预测时每个样本遍历全部树从根节点按x.get(splitFeat) splitValue走向左右子树最终把所有叶节点预测值乘以学习率累加见GBDTModel.predict()。对于希望快速本地验证的读者可直接运行仓库中的 GBDT 单元测试 GBDTTest.java它以 LOCAL 模式在单机单 Worker 单 PS 上完成训练与预测的端到端验证数据使用仓库自带的 agaricus 蘑菇数据集。4. 性能评测Angel vs XGBoost以下评测数据来自原文档基于腾讯内部数据集在腾讯线上 Gaia 集群YARN上完成对比。训练数据数据集数据集大小数据数量特征数量任务UserGender124GB1250 万2570二分类UserGender2145GB1.2 亿33 万二分类实验目的是预测用户的性别。数据集 UserGender1 大小为 24GB包含 1250 万个训练数据每个训练数据的特征维度为 2570数据集 UserGender2 大小为 145GB包含 1.2 亿个训练数据每个训练数据的特征维度为 33 万。两个数据集都是高维稀疏数据集。实验环境实验所使用的集群是腾讯的线上 Gaia 集群YARN单台机器的配置为CPU2680 × 2内存256 GB网络10G × 2磁盘4T × 12SATA参数配置Angel 和 XGBoost 使用如下的参数配置树的数量20树的最大高度7梯度直方图大小10学习速度0.1XGBoost、0.2Angel工作节点数量50参数服务器数量10每个工作节点内存2GBUserGender1、10GBUserGender2实验结果系统数据集训练总时间每棵树时间测试集误差XGBoostUserGender136min 48s110s0.155008AngelUserGender125min 22s76s0.154160XGBoostUserGender22h 25min435s0.232039AngelUserGender258min 39s175s0.243316需要说明的是上述结果来自该算法文档记录的特定评测实验数据集、集群与参数均如上用于说明 Angel 参数服务器架构在高维大数据集上的扩展性优势读者在自己的环境复现时应以实测为准。5. 总结与进一步阅读GBDT 的训练过程天然包含大模型存储 高频梯度同步 并行分裂搜索这与 Angel 参数服务器架构高度契合。通过将梯度直方图按 PS 节点切分存储、在两阶段分裂模式下让 PS 端并行计算局部最佳分裂点Angel 有效消除了高维特征下直方图聚合与分裂搜索的瓶颈实现了对 Spark 版 GBDT 与 MPI 版 XGBoost 的性能超越。如需进一步深入建议阅读仓库内的以下资源算法官方文档gbdt_on_angel.md、gbdt_on_angel_en.md分布式实现核心GBDTController.java、GBDTModel.scala任务入口与 RunnerGBDTRunner.scala、GBDTTrainTask.scala、GBDTPredictTask.scala配置定义MLConf.scala、GBDTParam.java单元测试与本地复现GBDTTest.java、GBDTLocalExample.java数据格式data_format.md参考文献Jiawei Jiang, Bin Cui, Ce Zhang and Fangcheng Fu. DimBoost: Boosting Gradient Boosting Decision Tree to Higher Dimensions. SIGMOD, 2018.Tianqi Chen and Carlos Guestrin. XGBoost: A Scalable Tree Boosting System. KDD, 2016.Michael Greenwald and Sanjeev Khanna. Space-efficient Online Computation of Quantile Summaries. SIGMOD, 2001.赞分享人工智能机器学习分布式训练图计算后端【免费下载链接】angelA Flexible and Powerful Parameter Server for large-scale machine learning项目地址https://gitcode.com/gh_mirrors/an/angel点击查看免费下载相关推荐MMPose全身姿态估计133关键点一次讲清怎么用MMPose全身姿态估计133关键点一次讲清怎么用 标准人体姿态估计只给 17 个关节指尖、五官、脚趾一概看不到。MMPose 的 WholeBody 全身计算机视觉人工智能深度学习3分钟掌握cargo-zigbuildRust跨平台编译的完整指南3分钟掌握cargo zigbuildRust跨平台编译的完整指南 cargo zigbuild是一个革命性的Rust构建工具它通过集成Zig编译器作为链接开发工具TensorFlow-Examples中的梯度提升决策树(GBDT)实现解析TensorFlow Examples中的梯度提升决策树 GBDT 实现解析 什么是梯度提升决策树 GBDT 梯度提升决策树 Gradient Boosted示例工程机器学习上一篇llama.cpp 代码评审技能PR 前自检清单与高频审查陷阱全解skills/code-review下一篇gsd-core 为何刻意不为 PLAN.md 提供“人类可读渲染”机器工件约定的边界与代价创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表