LRClassification.scala 1.4 KB

12345678910111213141516171819202122232425262728293031323334353637
  1. package com.example.model.classification
  2. import java.util
  3. import com.example.SparkConnect
  4. import com.example.model.helper.Utils
  5. import org.apache.spark.ml.{Pipeline, PipelineModel, PipelineStage}
  6. import org.apache.spark.ml.classification.{BinaryLogisticRegressionSummary, DecisionTreeClassificationModel, GBTClassifier, LogisticRegression, LogisticRegressionModel}
  7. import org.apache.spark.ml.evaluation.MulticlassClassificationEvaluator
  8. import org.apache.spark.ml.feature.{IndexToString, StringIndexer, VectorIndexer}
  9. import org.apache.spark.sql.{Dataset, Row}
  10. import org.springframework.beans.factory.annotation.Autowired
  11. import org.springframework.stereotype.Component
  12. import collection.JavaConverters._
  13. import com.example.util.{Constant, TrainParamEnum, TrainParamReader}
  14. @Component
  15. class LRClassification extends DefaultClassification {
  16. override def buildClassifier(params: util.Map[String, Object]): PipelineStage = {
  17. val maxIter = TrainParamReader.readParam(params,TrainParamEnum.MAXI_ITER).asInstanceOf[Int]
  18. val regParam = TrainParamReader.readParam(params,TrainParamEnum.REG_PARAM).asInstanceOf[Double]
  19. val elasticNetParam = TrainParamReader.readParam(params,TrainParamEnum.ELASTIC_NET_PARAM).asInstanceOf[Double]
  20. val lr = new LogisticRegression()
  21. .setLabelCol("label")
  22. .setFeaturesCol("indexedFeatures")
  23. .setMaxIter(maxIter)
  24. .setRegParam(regParam)
  25. .setElasticNetParam(elasticNetParam)
  26. lr
  27. }
  28. }