fix rabitEnv
This commit is contained in:
parent
808e30f9fc
commit
50337d1906
@ -17,6 +17,7 @@
|
|||||||
package ml.dmlc.xgboost4j.scala.spark
|
package ml.dmlc.xgboost4j.scala.spark
|
||||||
|
|
||||||
import scala.collection.mutable
|
import scala.collection.mutable
|
||||||
|
import scala.collection.JavaConverters._
|
||||||
|
|
||||||
import org.apache.commons.logging.LogFactory
|
import org.apache.commons.logging.LogFactory
|
||||||
import org.apache.spark.TaskContext
|
import org.apache.spark.TaskContext
|
||||||
@ -38,13 +39,13 @@ object XGBoost extends Serializable {
|
|||||||
private[spark] def buildDistributedBoosters(
|
private[spark] def buildDistributedBoosters(
|
||||||
trainingData: RDD[LabeledPoint],
|
trainingData: RDD[LabeledPoint],
|
||||||
xgBoostConfMap: Map[String, AnyRef],
|
xgBoostConfMap: Map[String, AnyRef],
|
||||||
|
rabitEnv: mutable.Map[String, String],
|
||||||
numWorkers: Int, round: Int, obj: ObjectiveTrait, eval: EvalTrait): RDD[Booster] = {
|
numWorkers: Int, round: Int, obj: ObjectiveTrait, eval: EvalTrait): RDD[Booster] = {
|
||||||
import DataUtils._
|
import DataUtils._
|
||||||
trainingData.repartition(numWorkers).mapPartitions {
|
trainingData.repartition(numWorkers).mapPartitions {
|
||||||
trainingSamples =>
|
trainingSamples =>
|
||||||
Rabit.init(new java.util.HashMap[String, String]() {
|
rabitEnv.put("DMLC_TASK_ID", TaskContext.getPartitionId().toString)
|
||||||
put("DMLC_TASK_ID", TaskContext.getPartitionId().toString)
|
Rabit.init(rabitEnv.asJava)
|
||||||
})
|
|
||||||
val dMatrix = new DMatrix(new JDMatrix(trainingSamples, null))
|
val dMatrix = new DMatrix(new JDMatrix(trainingSamples, null))
|
||||||
val booster = SXGBoost.train(xgBoostConfMap, dMatrix, round,
|
val booster = SXGBoost.train(xgBoostConfMap, dMatrix, round,
|
||||||
watches = new mutable.HashMap[String, DMatrix]{put("train", dMatrix)}.toMap, obj, eval)
|
watches = new mutable.HashMap[String, DMatrix]{put("train", dMatrix)}.toMap, obj, eval)
|
||||||
@ -59,7 +60,8 @@ object XGBoost extends Serializable {
|
|||||||
val sc = trainingData.sparkContext
|
val sc = trainingData.sparkContext
|
||||||
val tracker = new RabitTracker(numWorkers)
|
val tracker = new RabitTracker(numWorkers)
|
||||||
require(tracker.start(), "FAULT: Failed to start tracker")
|
require(tracker.start(), "FAULT: Failed to start tracker")
|
||||||
boosters = buildDistributedBoosters(trainingData, configMap, numWorkers, round, obj, eval)
|
boosters = buildDistributedBoosters(trainingData, configMap,
|
||||||
|
tracker.getWorkerEnvs.asScala, numWorkers, round, obj, eval)
|
||||||
// force the job
|
// force the job
|
||||||
boosters.foreachPartition(_ => ())
|
boosters.foreachPartition(_ => ())
|
||||||
println("=====finished training=====")
|
println("=====finished training=====")
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user