- 人工智能
- 机器学习
【免费下载链接】xgboost
Scalable, Portable and Distributed Gradient Boosting (GBDT, GBRT or GBM) Library, for Python, R, Java, Scala, C and more. Runs on single machine, Hadoop, Spark, Dask, Flink and DataFlow
本篇技术指南以仓库中jvm-packages/xgboost4j-example/README.md为核心骨架,系统梳理 XGBoost4J(XGBoost 的 JVM 生态)在 Java、Scala、Spark、Flink 四种 API 下的官方示例代码。你将掌握:如何加载 LIBSVM/CSV 数据构造 DMatrix、训练与预测、自定义损失函数与评估指标、从已有预测结果继续提升(boosting from prediction)、仅用前 N 棵树预测、预测叶子索引、广义线性模型(gblinear)、交叉验证、早停机制,以及基于 Spark MLlib Pipeline 和 Flink DataSet 的分布式训练完整流程。
一、示例工程总览:四大 API 一览
官方示例工程位于仓库的 jvm-packages/xgboost4j-example,配套的 Maven 工程定义见 pom.xml。示例按编程接口划分为四大类:
| API | 语言 | 核心示例 |
|---|---|---|
| Java API | Java | 基础流程、自定义损失/评估、从预测继续提升、前 N 棵树预测、广义线性模型、交叉验证、叶子索引预测、早停 |
| Scala API | Scala | 与 Java 一一对应的同套示例 |
| Spark API | Scala | 基于 Spark MLlib Pipeline 的分布式训练 |
| Flink API | Scala / Java | 基于 Flink DataSet 的分布式训练 |
示例目录结构如下(源码位于src/main,测试位于src/test):
jvm-packages/xgboost4j-example/src/main/ ├── java/ml/dmlc/xgboost4j/java/example/ │ ├── BasicWalkThrough.java # 基础流程(训练/预测/保存/加载) │ ├── CustomObjective.java # 自定义损失函数与评估指标 │ ├── BoostFromPrediction.java # 从已有预测结果继续提升 │ ├── PredictFirstNtree.java # 仅用前 N 棵树预测 │ ├── GeneralizedLinearModel.java # gblinear 广义线性模型 │ ├── CrossValidation.java # 交叉验证 │ ├── PredictLeafIndices.java # 预测叶子索引 │ ├── EarlyStopping.java # 早停机制 │ └── util/{DataLoader, CustomEval}.java └── scala/ml/dmlc/xgboost4j/scala/example/ ├── BasicWalkThrough.scala # 对应 Java 版 ├── spark/SparkMLlibPipeline.scala # Spark 分布式训练 └── flink/DistTrainWithFlink.scala # Flink 分布式训练这些示例全部有对应的自动化测试保障:Java 示例由 JavaExamplesTest.java 覆盖,Scala 示例由 ScalaExamplesTest.scala 与 SparkExamplesTest.scala 覆盖,Flink 示例由 DistTrainWithFlinkSuite.scala 覆盖,读者可结合测试用例验证示例的可运行性。
示例默认使用的数据是仓库demo/data/下的蘑菇分类数据集(LIBSVM 格式):训练集agaricus.txt.train、测试集agaricus.txt.test,以及特征名映射文件featmap.txt(可参考 demo/data 下的 README 了解数据来源)。
二、Java API:基础流程(BasicWalkThrough)
示例文件:BasicWalkThrough.java
2.1 加载数据:构造 DMatrix
DMatrix 是 XGBoost 的数据结构抽象,示例演示了三种构造方式:
// 1. 直接以文本文件路径构造(LIBSVM 格式,? 后跟解析参数) DMatrix trainMat = new DMatrix( "../../demo/data/agaricus.txt.train?format=libsvm&indexing_mode=1"); DMatrix testMat = new DMatrix( "../../demo/data/agaricus.txt.test?format=libsvm&indexing_mode=1");其中indexing_mode=1表示特征索引从 1 开始计数(LIBSVM 惯例),这在特征索引与featmap.txt对齐时很重要。
第二种构造方式是从 CSR 稀疏矩阵构造(见BasicWalkThrough.java#L115-L122):
// 先用 DataLoader 解析 LIBSVM 文件为 CSR 三段式结构 DataLoader.CSRSparseData spData = DataLoader.loadSVMFile("../../demo/data/agaricus.txt.train"); // rowHeaders: 每行起始偏移;colIndex: 每列索引;data: 非零值;127 为特征总数 DMatrix trainMat2 = new DMatrix(spData.rowHeaders, spData.colIndex, spData.data, DMatrix.SparseType.CSR, 127); trainMat2.setLabel(spData.labels); // 设置标签DataLoader是示例自带的工具类(DataLoader.java),提供loadSVMFile(解析label idx:value格式)与loadCSVFile(解析 CSV,最后一列为标签)两种解析器,返回CSRSparseData/DenseData结构体。
2.2 训练参数、watchlist 与训练
HashMap<String, Object> params = new HashMap<String, Object>(); params.put("eta", 1.0); // 学习率(步长收缩) params.put("max_depth", 2); // 树最大深度 params.put("silent", 1); // 静默模式,不打印训练日志 params.put("objective", "binary:logistic"); // 二分类逻辑回归目标函数 HashMap<String, DMatrix> watches = new HashMap<String, DMatrix>(); watches.put("train", trainMat); // 训练集(用于打印训练评估) watches.put("test", testMat); // 测试集(用于监控泛化表现) int round = 2; // boosting 迭代轮数 Booster booster = XGBoost.train(trainMat, params, round, watches, null, null);XGBoost.train的完整签名为train(DMatrix dtrain, Map params, int round, Map watches, IObjective obj, IEvaluation eval),后两个参数为自定义目标函数与评估函数,传null表示使用内置实现。
2.3 预测、保存与加载模型
float[][] predicts = booster.predict(testMat); // 预测 booster.saveModel("./model/xgb.model"); // 保存模型 String[] modelInfos = booster.getModelDump("../../demo/data/featmap.txt", false); saveDumpModel("./model/dump.raw.txt", modelInfos); // 转储模型为文本(含特征名) testMat.saveBinary("./model/dtest.buffer"); // DMatrix 保存为二进制缓存 Booster booster2 = XGBoost.loadModel("./model/xgb.model"); // 重新加载模型 DMatrix testMat2 = new DMatrix("./model/dtest.buffer"); // 加载二进制数据 float[][] predicts2 = booster2.predict(testMat2); System.out.println(checkPredicts(predicts, predicts2)); // 验证预测结果一致示例用checkPredicts方法逐行比较两次预测结果,证明「保存→加载」链路无损。getModelDump配合featmap.txt可输出带特征名的树结构文本,用于模型可解释性分析。
三、Java API:自定义损失函数与评估指标(CustomObjective)
示例文件:CustomObjective.java
XGBoost4J 允许用户实现IObjective(自定义损失)与IEvaluation(自定义评估指标)两个接口,并注入XGBoost.train。
3.1 实现自定义目标函数
以逻辑回归损失为例,需要实现IObjective.getGradient返回一阶导数(grad)与二阶导数(hess):
public static class LogRegObj implements IObjective { public float sigmoid(float input) { return (float) (1 / (1 + Math.exp(-input))); // 示例实现,未做数值稳定处理 } @Override public List<float[]> getGradient(float[][] predicts, DMatrix dtrain) { int nrow = predicts.length; float[] labels = dtrain.getLabel(); float[] grad = new float[nrow]; // 一阶梯度 float[] hess = new float[nrow]; // 二阶梯度(Hessian) float[][] transPredicts = transform(predicts); for (int i = 0; i < nrow; i++) { float predict = transPredicts[i][0]; // sigmoid 转换后的概率 grad[i] = predict - labels[i]; // 逻辑回归梯度 hess[i] = predict * (1 - predict); // 逻辑回归 Hessian } gradients.add(grad); gradients.add(hess); return gradients; } }3.2 实现自定义评估指标
public static class EvalError implements IEvaluation { @Override public String getMetric() { return "custom_error"; } // 指标名,将出现在训练日志中 @Override public float eval(float[][] predicts, DMatrix dmat) { float error = 0f; float[] labels = dmat.getLabel(); for (int i = 0; i < nrow; i++) { if (labels[i] == 0f && predicts[i][0] > 0) error++; // 误判为正 else if (labels[i] == 1f && predicts[i][0] <= 0) error++; // 误判为负 } return error / labels.length; } }3.3 注入自定义实现并训练
IObjective obj = new LogRegObj(); IEvaluation eval = new EvalError(); Booster booster = XGBoost.train(trainMat, params, round, watches, obj, eval);重要注意事项(源码注释明确强调):自定义损失函数时,默认预测输出是 margin(即 sigmoid 变换前的原始得分)。这会导致内置评估指标失效——例如内置的error指标假设输入已经是 sigmoid 后的概率,而实际拿到的是原始得分。因此自定义损失后通常需要同时自定义评估函数。
四、Java API:从已有预测继续提升(BoostFromPrediction)
示例文件:BoostFromPrediction.java
该示例演示了 XGBoost 的base_margin机制——将一个已训练模型的预测结果作为下一轮训练的初始预测(bias),实现"接力训练"或增量提升:
// 先训练 1 轮得到初始模型 Booster booster = XGBoost.train(trainMat, params, 1, watches, null, null); // 输出 margin 形式的预测(第二个参数 true 表示输出原始 margin) float[][] trainPred = booster.predict(trainMat, true); float[][] testPred = booster.predict(testMat, true); // 将预测结果设置为 DMatrix 的 base margin trainMat.setBaseMargin(trainPred); testMat.setBaseMargin(testPred); // 基于初始预测继续训练 Booster booster2 = XGBoost.train(trainMat, params, 1, watches, null, null);注意predict(trainMat, true)的第二个布尔参数:true表示输出未经 sigmoid/softmax 变换的 margin 值。这在多模型集成、增量学习、以及需要精确控制初始得分的场景(如推荐系统的冷启动打分)中非常实用。
五、Java API:前 N 棵树预测与叶子索引预测
5.1 仅用前 N 棵树预测(PredictFirstNtree)
示例文件:PredictFirstNtree.java
int round = 3; Booster booster = XGBoost.train(trainMat, params, round, watches, null, null); // 只使用前 1 棵树的预测 float[][] predicts1 = booster.predict(testMat, false, 1); // 默认使用全部 boosting 迭代的预测 float[][] predicts2 = booster.predict(testMat); CustomEval eval = new CustomEval(); System.out.println("error of predicts1: " + eval.eval(predicts1, testMat)); System.out.println("error of predicts2: " + eval.eval(predicts2, testMat));predict(DMatrix, boolean outputMargin, int ntreeLimit)的第三个参数ntreeLimit控制参与预测的树数量上限。该能力在以下场景有价值:
- 模型压缩 / 推理加速:按精度要求截断树数量;
- 观察模型在不同迭代轮次的预测质量曲线,辅助选择早停轮次;
- 集成学习中的"部分模型"复用。
5.2 预测叶子索引(PredictLeafIndices)
示例文件:PredictLeafIndices.java
int round = 3; Booster booster = XGBoost.train(trainMat, params, round, watches, null, null); // 只使用前 2 轮迭代的叶子索引 float[][] leafindex = booster.predictLeaf(testMat, 2); // 使用全部迭代的叶子索引(ntreeLimit = 0 表示全部) leafindex = booster.predictLeaf(testMat, 0); System.out.println(leafindex[0][0] + ", " + leafindex[0][1]);predictLeaf返回每个样本在每棵树中所落叶子节点的索引。leaf index 本质上是将原始特征空间映射为高维稀疏的类别特征,常用于:
- 将 XGBoost 作为特征转换器,输出叶子索引后接入 LR / FM 等模型(经典的 GBDT+LR 方案);
- 分析样本在树结构中的分布,做异常检测或可解释性分析。
六、Java API:广义线性模型与交叉验证
6.1 广义线性模型(GeneralizedLinearModel)
示例文件:GeneralizedLinearModel.java
将booster参数设为gblinear即可从树模型切换到线性模型:
HashMap<String, Object> params = new HashMap<String, Object>(); params.put("alpha", 0.0001); // L1 正则化系数 params.put("silent", 1); params.put("objective", "binary:logistic"); params.put("booster", "gblinear"); // 关键:线性基学习器 int round = 4; Booster booster = XGBoost.train(trainMat, params, round, watches, null, null); float[][] predicts = booster.predict(testMat);示例注释中给出两点重要提示:
lambda为 L2 正则化系数,还可设置lambda_bias(偏置项的 L2 正则);- XGBoost 的 gblinear 使用并行坐标下降算法(shotgun),并行化在特定情况下可能影响收敛;通常无需显式设置
eta(步长),但将eta调小(如 0.5)可使优化更稳定。
6.2 交叉验证(CrossValidation)
示例文件:CrossValidation.java
DMatrix trainMat = new DMatrix("../../demo/data/agaricus.txt.train?format=libsvm"); HashMap<String, Object> params = new HashMap<String, Object>(); params.put("eta", 1.0); params.put("max_depth", 3); params.put("silent", 1); params.put("nthread", 6); // 线程数 params.put("objective", "binary:logistic"); params.put("gamma", 1.0); // 分裂增益阈值(最小损失减少) params.put("eval_metric", "error");// 评估指标 int round = 2; // boosting 轮数 int nfold = 5; // 5 折交叉验证 String[] metrics = null; // 可传入附加评估指标数组 String[] evalHist = XGBoost.crossValidation(trainMat, params, round, nfold, metrics, null, null);XGBoost.crossValidation返回各轮各折的评估历史字符串数组,用于观察模型在不同折上的稳定性与方差。
七、Java API:早停机制(EarlyStopping)
示例文件:EarlyStopping.java
该示例演示了XGBoost.train重载形式中带早停参数的版本:
Map<String, Object> paramMap = new HashMap<String, Object>() {{ put("max_depth", 3); put("objective", "binary:logistic"); put("maximize_evaluation_metrics", "false"); // 指标越小越好(如 error) }}; int nRounds = 128; // 最大迭代轮数 int nEarlyStoppingRounds = 4; // 连续 n 轮指标无改善即停止 // 注意:本示例将评估历史通过 metrics 二维数组传出([数据集][轮次]) float[][] metrics = new float[watches.size()][nRounds]; Booster booster = XGBoost.train(trainXy, paramMap, nRounds, watches, metrics, null, null, nEarlyStoppingRounds); // 早停后通过 Booster 属性读取最佳迭代与最佳得分 int bestIter = Integer.valueOf(booster.getAttr("best_iteration")); float bestScore = Float.valueOf(booster.getAttr("best_score")); System.out.printf("Best iter: %d, Best score: %f\n", bestIter, bestScore);关键技术点:
- 早停前必须通过
maximize_evaluation_metrics明确指标优化方向(true为越大越好,如 auc;false为越小越好,如 error); - 早停后的最佳结果保存在 Booster 的
best_iteration与best_score属性中,可通过booster.getAttr(...)读取; - 相比固定轮数训练,早停可避免过拟合并显著节省训练时间。
八、Scala API:与 Java 一一对应的流畅体验
Scala 示例位于src/main/scala/ml/dmlc/xgboost4j/scala/example/,与 Java 示例一一对应,包括BasicWalkThrough.scala、CustomObjective.scala、BoostFromPrediction.scala、PredictFirstNTree.scala、GeneralizedLinearModel.scala、CrossValidation.scala、PredictLeafIndices.scala,并复用 Java 侧的DataLoader工具类。
以 BasicWalkThrough.scala 为例,可以看到 Scala 版在 API 上更为简洁:
import ml.dmlc.xgboost4j.scala.{DMatrix, XGBoost} import scala.collection.mutable val trainMax = new DMatrix("../../demo/data/agaricus.txt.train?format=libsvm&indexing_mode=1") val params = new mutable.HashMap[String, Any]() params += "eta" -> 1.0 params += "max_depth" -> 2 params += "silent" -> 1 params += "objective" -> "binary:logistic" val watches = new mutable.HashMap[String, DMatrix] watches += "train" -> trainMax watches += "test" -> testMax val booster = XGBoost.train(trainMax, params.toMap, round, watches.toMap) val predicts = booster.predict(testMax) booster.saveModel(file.getAbsolutePath + "/xgb.model") val modelInfos = booster.getModelDump("../../demo/data/featmap.txt", false) saveDumpModel(file.getAbsolutePath + "/dump.raw.txt", modelInfos) testMax.saveBinary(file.getAbsolutePath + "/dtest.buffer") val booster2 = XGBoost.loadModel(file.getAbsolutePath + "/xgb.model") val testMax2 = new DMatrix(file.getAbsolutePath + "/dtest.buffer") val predicts2 = booster2.predict(testMax2) println(checkPredicts(predicts, predicts2))Scala 封装位于jvm-packages/xgboost4j模块内(如ml.dmlc.xgboost4j.scala.XGBoost),在 Java API 之上提供更符合 Scala 惯用法的集合类型与函数式调用方式。Scala 用户可以参考 ScalaExamplesTest.scala 中的测试来验证每个示例的预期行为。
九、Spark API:基于 MLlib Pipeline 的分布式训练
示例文件:SparkMLlibPipeline.scala
该示例以经典 Iris(鸢尾花)数据集演示 XGBoost 与 Spark MLlib Pipeline 的完整集成,命令用法为:
SparkMLlibPipeline input_path native_model_path pipeline_model_path [cpu|gpu]最后一个参数可选cpu(默认,2 个 worker)或gpu(device=cuda,1 个 worker)。
9.1 读取数据并划分训练 / 测试集
val schema = new StructType(Array( StructField("sepal length", DoubleType, true), StructField("sepal width", DoubleType, true), StructField("petal length", DoubleType, true), StructField("petal width", DoubleType, true), StructField("class", StringType, true))) val rawInput = spark.read.schema(schema).csv(inputPath) val Array(training, test) = rawInput.randomSplit(Array(0.8, 0.2), 123)9.2 构建四阶段 Pipeline
Pipeline 共四个阶段:
VectorAssembler:将 4 个特征列组装为单个向量列features;StringIndexer:将字符串标签class转为数值索引classIndex;XGBoostClassifier:XGBoost 分类器(核心训练器);IndexToString:将预测索引转回原始字符串标签。
import ml.dmlc.xgboost4j.scala.spark.{XGBoostClassificationModel, XGBoostClassifier} val assembler = new VectorAssembler() .setInputCols(Array("sepal length", "sepal width", "petal length", "petal width")) .setOutputCol("features") val labelIndexer = new StringIndexer() .setInputCol("class").setOutputCol("classIndex").fit(training) val booster = new XGBoostClassifier(Map( "eta" -> 0.1f, "max_depth" -> 2, "objective" -> "multi:softprob", // 多分类 softmax 概率输出 "num_class" -> 3, // 类别数 "device" -> device // "cpu" 或 "cuda" )).setNumRound(10).setNumWorkers(numWorkers) booster.setFeaturesCol("features") booster.setLabelCol("classIndex") val labelConverter = new IndexToString() .setInputCol("prediction").setOutputCol("realLabel") .setLabels(labelIndexer.labelsArray(0)) val pipeline = new Pipeline().setStages(Array(assembler, labelIndexer, booster, labelConverter)) val model: PipelineModel = pipeline.fit(training)9.3 批量预测、评估与超参数调优
val prediction = model.transform(test) prediction.show(false) val evaluator = new MulticlassClassificationEvaluator() evaluator.setLabelCol("classIndex").setPredictionCol("prediction") println("The model accuracy is : " + evaluator.evaluate(prediction)) // 使用 Spark CrossValidator 做网格搜索:max_depth ∈ {3, 8},eta ∈ {0.2, 0.6} val paramGrid = new ParamGridBuilder() .addGrid(booster.maxDepth, Array(3, 8)) .addGrid(booster.eta, Array(0.2, 0.6)) .build() val cv = new CrossValidator() .setEstimator(pipeline).setEvaluator(evaluator) .setEstimatorParamMaps(paramGrid).setNumFolds(3) val cvModel = cv.fit(training) val bestModel = cvModel.bestModel.asInstanceOf[PipelineModel].stages(2) .asInstanceOf[XGBoostClassificationModel] println("best model params: " + bestModel.extractParamMap())9.4 模型导出与 Pipeline 持久化
// 导出为原生 XGBoost 模型,可在本地 Python 环境加载 bestModel.nativeBooster.saveModel(nativeModelPath) // Pipeline 模型持久化与加载 model.write.overwrite().save(pipelineModelPath) val model2 = PipelineModel.load(pipelineModelPath) model2.transform(test)这一节展示了 XGBoost4J-Spark 的三个关键卖点:完全兼容 Spark MLlib Pipeline / CrossValidator / ParamGridBuilder 生态、通过nativeBooster与原生 XGBoost 模型格式互通(训练产物可直接跨语言复用)、gpu 模式仅需切换device参数。Spark 相关的更多示例还可参考同目录下的 SparkTraining.scala。
十、Flink API:基于 DataSet 的分布式训练
10.1 Java 版(DistTrainWithFlinkExample)
示例文件:DistTrainWithFlinkExample.java
Flink 示例基于 Flink ML 的Vector抽象与 DataSet API,演示了 CSV 数据读取、按比例切分训练 / 测试集、分布式训练与预测的完整链路:
import ml.dmlc.xgboost4j.java.flink.XGBoost; import ml.dmlc.xgboost4j.java.flink.XGBoostModel; // 读取 CSV 并打包为 (Vector, Double) 数据集 final DataSet<Tuple2<Vector, Double>> trainData = ...; // 前 percentage% 作为训练集 final DataSet<Vector> testData = ...; // 其余作为测试集 // 定义参数 HashMap<String, Object> paramMap = new HashMap<String, Object>(3); paramMap.put("eta", 0.1); paramMap.put("max_depth", 2); paramMap.put("objective", "binary:logistic"); // 训练与预测(round = 2) XGBoostModel model = XGBoost.train(trainData, paramMap, round); DataSet<Float[]> predTest = model.predict(testData); List<Float[]> list = predTest.collect();示例配套数据为veterans_lung_cancer.csv(退伍军人肺癌数据集),通过DataSetUtils.zipWithIndex为数据打上序号后按前 70% / 后 30% 切分。mapFunction将 CSV 行转为 Flink 的DenseVector,并根据字符串标签是否包含"inf"映射为二分类标签。
10.2 Scala 版与测试
Scala 版见 DistTrainWithFlink.scala,其训练逻辑与 Java 版一致。对应的测试套件 DistTrainWithFlinkSuite.scala 与 DistTrainWithFlinkExampleTest.scala 可用于验证分布式训练链路是否正常工作。
十一、如何运行这些示例
11.1 数据准备
示例默认使用仓库demo/data/下的agaricus.txt.train与agaricus.txt.test(LIBSVM 格式)。在运行示例前需确认:
- 相对路径
../../demo/data/是从jvm-packages/xgboost4j-example模块的工作目录出发的,若从仓库根目录运行,请改用demo/data/agaricus.txt.train之类的仓库根相对路径; - Flink 示例需要准备
veterans_lung_cancer.csv数据文件,Spark 示例需要 Iris 格式的 CSV 数据。
11.2 构建与执行
jvm-packages是 Maven 多模块工程,根 pom.xml 聚合了xgboost4j、xgboost4j-example、xgboost4j-spark、xgboost4j-flink等模块,xgboost4j-example自身的构建配置见其 pom.xml。典型流程:
# 在 jvm-packages 目录下构建(需先完成原生库编译,具体可参考 doc/jvm 目录下的文档) mvn -pl xgboost4j-example -am package # 运行 Java 示例 mvn -pl xgboost4j-example exec:java -Dexec.mainClass=ml.dmlc.xgboost4j.java.example.BasicWalkThrough # 运行 Spark 示例(提交到 Spark 集群或本地模式) spark-submit --class ml.dmlc.xgboost4j.scala.example.spark.SparkMLlibPipeline \ xgboost4j-example-<version>.jar \ /path/to/iris.csv /path/to/native_model /path/to/pipeline_model cpuSpark 示例的命令行参数为input_path native_model_path pipeline_model_path [cpu|gpu],运行时通过randomSplit(Array(0.8, 0.2), 123)固定随机种子保证可复现性。
十二、从示例到原理:配套核心源码导读
示例背后对应的核心封装位于 jvm-packages/xgboost4j 模块(含 Java/Scala 双封装、JNI 桥接层xgboost4j/src/main/native等)。建议按以下路径深入源码:
| 学习目标 | 推荐源码路径 |
|---|---|
DMatrix 的三种构造方式与setBaseMargin/saveBinary实现 | xgboost4j模块中ml.dmlc.xgboost4j.java.DMatrix |
XGBoost.train/crossValidation/loadModel的 JNI 调用链 | xgboost4j模块中ml.dmlc.xgboost4j.java.XGBoost及XGBoostJNI |
IObjective/IEvaluation接口的 JNI 回调机制 | 同上模块中IObjective.java/IEvaluation.java及对应 native 代码 |
predict(..., ntreeLimit)/predictLeaf的底层语义 | 原生 C++ 侧src/predictor/目录(预测器实现) |
Spark 封装XGBoostClassifier/XGBoostClassificationModel | jvm-packages/xgboost4j-spark/src/main/scala/ |
Flink 封装XGBoost.train(DataSet, ...) | jvm-packages/xgboost4j-flink/src/main/ |
例如,booster.getAttr("best_iteration")读取早停最佳轮次的能力,对应 Booster 内部维护的模型属性表;predictLeaf输出叶节点索引的能力,对应原生预测器中 leaf index 预测模式。理解这些对应关系后,读者可以在此基础上按需扩展自己的业务封装。
结语
本工程示例覆盖了 XGBoost4J 从单机 Java/Scala 到分布式 Spark/Flink 的全部主流用法,且每个示例均有对应测试保障。建议学习路径:先用BasicWalkThrough打通「数据 → 训练 → 预测 → 保存/加载」的完整闭环,再按业务需求依次实践自定义损失、早停、叶子索引等进阶能力,最后通过 Spark / Flink 示例将模型落地到分布式生产环境。更多 JVM 生态的工程化细节(如依赖管理、原生库打包、GPU 支持)可参阅仓库 doc/jvm 下的文档。
- 人工智能
- 机器学习
【免费下载链接】xgboost
Scalable, Portable and Distributed Gradient Boosting (GBDT, GBRT or GBM) Library, for Python, R, Java, Scala, C and more. Runs on single machine, Hadoop, Spark, Dask, Flink and DataFlow
相关推荐
3步构建专业量化交易系统:Lean引擎实战指南
3步构建专业量化交易系统:Lean引擎实战指南 当量化交易从华尔街的专属技能转变为每个开发者都能触及的技术时,如何快速搭建一个既专业又灵活的交易系统?Quant
金融科技后端Delta Lake 官方示例全解析:从 Quickstart 到流处理、CDF 与 UniForm 的 Scala/Python 实战指南
Delta Lake 官方示例全解析:从 Quickstart 到流处理、CDF 与 UniForm 的 Scala/Python 实战指南 Delta Lak
湖仓一体数据工程数据湖超强DocsGPT API实战指南:接口全解析+示例代码
超强DocsGPT API实战指南:接口全解析+示例代码 你还在为API调用复杂、文档不清晰而烦恼吗?想快速掌握DocsGPT的强大功能,却被繁琐的技术细节困住
人工智能AI 应用AI AgentRAG后端前端MCP 服务深度研究
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考