dmlc / dmlc/xgboost

[jvm-packages] Does xgboost4j 1.3+ support survival models?

Open
#8,670 4 comments 0 reactions 0 assignees View on GitHub
feature-request
Dominant language
C++
Stars
28.8k
Forks
8.9k
Avg merge
1d 12h
Merged PRs (30d)
54

Description

Environment: xgboost4j_2.12:1.3.1 / spark 3.+

We'd like to use xgboost4j with the aft objective and could not find any existing examples.
I adapted the codes in https://xgboost.readthedocs.io/en/stable/tutorials/aft_survival_analysis.html and got the following error. Shall I represent labels differently? we also tried xgboost4j 1.7.1 and got the same error.

import org.apache.spark.ml.linalg.Vectors.dense
val dataval = spark.createDataFrame(Seq(
(dense(1.0 ,-1.0), dense(2.0,2.0)),
(dense(-1.0, 1.0), dense(3.0,Integer.MAX_VALUE)),
(dense(0.0 , 1.0), dense(0.0,4.0)),
(dense(1.0 , 0.0), dense(4.0,5.0))
)).toDF("features", "label")

dataval.show
+----------+-------------------+
| features| label|
+----------+-------------------+
|[1.0,-1.0]| [2.0,2.0]|
|[-1.0,1.0]|[3.0,2.147483647E9]|
| [0.0,1.0]| [0.0,4.0]|
| [1.0,0.0]| [4.0,5.0]|
+----------+-------------------+

// training parameters
val paramMap = List(
"objective" -> "survival:aft",
"eval_metric" -> "aft-nloglik",
"aft_loss_distribution" -> "normal",
"aft_loss_distribution_scale" -> 1.20,
"tree_method" -> "hist",
"learning_rate" -> 0.05,
"max_depth" -> 2,
"num_workers" -> 1,
"disable_default_eval_metric" -> true
).toMap

println("Starting Xgboost ")
val xgbClassifier = new XGBoostClassifier(paramMap).
setFeaturesCol("features").
setLabelCol("label")
val model = xgbClassifier.fit(dataval)

An error was encountered:
java.lang.IllegalArgumentException: requirement failed: Column label must be of type numeric but was actually of type struct,values:array>.
at scala.Predef$.require(Predef.scala:281)
at org.apache.spark.ml.util.SchemaUtils$.checkNumericType(SchemaUtils.scala:78)
at org.apache.spark.ml.PredictorParams.validateAndTransformSchema(Predictor.scala:54)
at org.apache.spark.ml.PredictorParams.validateAndTransformSchema$(Predictor.scala:47)
at org.apache.spark.ml.classification.Classifier.org$apache$spark$ml$classification$ClassifierParams$$super$validateAndTransformSchema(Classifier.scala:73)
at org.apache.spark.ml.classification.ClassifierParams.validateAndTransformSchema(Classifier.scala:43)
at org.apache.spark.ml.classification.ClassifierParams.validateAndTransformSchema$(Classifier.scala:39)
at org.apache.spark.ml.classification.ProbabilisticClassifier.org$apache$spark$ml$classification$ProbabilisticClassifierParams$$super$validateAndTransformSchema(ProbabilisticClassifier.scala:51)
at org.apache.spark.ml.classification.ProbabilisticClassifierParams.validateAndTransformSchema(ProbabilisticClassifier.scala:38)
at org.apache.spark.ml.classification.ProbabilisticClassifierParams.validateAndTransformSchema$(ProbabilisticClassifier.scala:34)
at org.apache.spark.ml.classification.ProbabilisticClassifier.validateAndTransformSchema(ProbabilisticClassifier.scala:51)
at org.apache.spark.ml.Predictor.transformSchema(Predictor.scala:177)
at org.apache.spark.ml.PipelineStage.transformSchema(Pipeline.scala:71)
at org.apache.spark.ml.Predictor.fit(Predictor.scala:133)
... 53 elided

Contributor guide

No contributing guide indexed for this repository

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.