news 2026/9/20 14:57:45

XGBoost4J 官方代码示例全解:Java / Scala / Spark / Flink 四大 API 实战指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
XGBoost4J 官方代码示例全解:Java / Scala / Spark / Flink 四大 API 实战指南
  • 人工智能
  • 机器学习

【免费下载链接】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

项目地址:https://gitcode.com/gh_mirrors/xg/xgboost
点击查看免费下载

本篇技术指南以仓库中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 APIJava基础流程、自定义损失/评估、从预测继续提升、前 N 棵树预测、广义线性模型、交叉验证、叶子索引预测、早停
Scala APIScala与 Java 一一对应的同套示例
Spark APIScala基于 Spark MLlib Pipeline 的分布式训练
Flink APIScala / 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_iterationbest_score属性中,可通过booster.getAttr(...)读取;
  • 相比固定轮数训练,早停可避免过拟合并显著节省训练时间。

八、Scala API:与 Java 一一对应的流畅体验

Scala 示例位于src/main/scala/ml/dmlc/xgboost4j/scala/example/,与 Java 示例一一对应,包括BasicWalkThrough.scalaCustomObjective.scalaBoostFromPrediction.scalaPredictFirstNTree.scalaGeneralizedLinearModel.scalaCrossValidation.scalaPredictLeafIndices.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)或gpudevice=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 共四个阶段:

  1. VectorAssembler:将 4 个特征列组装为单个向量列features
  2. StringIndexer:将字符串标签class转为数值索引classIndex
  3. XGBoostClassifier:XGBoost 分类器(核心训练器);
  4. 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.trainagaricus.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 聚合了xgboost4jxgboost4j-examplexgboost4j-sparkxgboost4j-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 cpu

Spark 示例的命令行参数为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.XGBoostXGBoostJNI
IObjective/IEvaluation接口的 JNI 回调机制同上模块中IObjective.java/IEvaluation.java及对应 native 代码
predict(..., ntreeLimit)/predictLeaf的底层语义原生 C++ 侧src/predictor/目录(预测器实现)
Spark 封装XGBoostClassifier/XGBoostClassificationModeljvm-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

项目地址:https://gitcode.com/gh_mirrors/xg/xgboost
点击查看免费下载

相关推荐

上一篇:Jellyfin Media Player硬件加速解码:释放你的GPU潜力
下一篇:Mihon漫画阅读器:免费开源的Android漫画阅读终极指南

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/20 14:57:34

JMeter压测实战:从脚本设计到性能瓶颈定位

1. 压测前必须想清楚的几件事做JMeter压测&#xff0c;最大的坑不是工具不会用&#xff0c;而是不知道为什么压、压什么、压多久。我见过太多团队拿到JMeter就咔咔加线程数&#xff0c;结果压出来的报告除了给自己壮胆&#xff0c;没有任何参考价值。如果手里正好有一个“迁移到…

作者头像 李华
网站建设 2026/9/20 14:57:04

Wireshark安装核心原理:Npcap驱动与架构匹配详解

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/20 14:54:55

电力工程二次设计必备:国际图集D部分的高效利用与资料管理实战

简介&#xff1a;这是一份面向电气设计与施工从业者的电力图集PDF资料&#xff0c;聚焦国际图集电力D部分&#xff0c;涵盖变配电二次接线、防雷接地、电缆敷设、低压配电与电动机控制、舞台灯光、住宅小区电气设计等关键环节&#xff0c;适用于民用建筑与工业环境中的方案参考…

作者头像 李华
网站建设 2026/9/20 14:54:20

零代码本地部署Gemini 3.1 Flash:Ollama+Open WebUI实战指南

1. 为什么要在本地跑一个 Flash 级模型先把结论摆在前面&#xff1a;如果你手头有一台内存 16GB 以上的普通笔记本&#xff0c;或者一台带独显的台式机&#xff0c;那么用 Ollama 加 Open WebUI 把 Gemini 3.1 Flash 这类轻量模型跑起来&#xff0c;是一件当天就能搞定的事。整…

作者头像 李华
网站建设 2026/9/20 14:53:36

LibreChat自托管AI对话平台:多模型统一入口的Docker部署实操指南

我大概花了两个晚上&#xff0c;把本地一直凑合用的几套AI聊天客户端全部换掉&#xff0c;统一指向了LibreChat。如果你最近也在GitHub上刷到过这个项目&#xff0c;或者正被“想用多个模型又不想开一堆网页标签”的问题困扰&#xff0c;那这篇实操笔记应该能帮你省下不少弯路。…

作者头像 李华
网站建设 2026/9/20 14:53:21

USB 3.2 Gen2x2真相:20Gbps为何跑不满?

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华