| 12345678910111213141516171819202122232425262728293031323334353637 |
- package com.example.model.classification
- import java.util
- import com.example.SparkConnect
- import com.example.model.helper.Utils
- import org.apache.spark.ml.{Pipeline, PipelineModel, PipelineStage}
- import org.apache.spark.ml.classification.{BinaryLogisticRegressionSummary, DecisionTreeClassificationModel, GBTClassifier, LogisticRegression, LogisticRegressionModel}
- import org.apache.spark.ml.evaluation.MulticlassClassificationEvaluator
- import org.apache.spark.ml.feature.{IndexToString, StringIndexer, VectorIndexer}
- import org.apache.spark.sql.{Dataset, Row}
- import org.springframework.beans.factory.annotation.Autowired
- import org.springframework.stereotype.Component
- import collection.JavaConverters._
- import com.example.util.{Constant, TrainParamEnum, TrainParamReader}
- @Component
- class LRClassification extends DefaultClassification {
- override def buildClassifier(params: util.Map[String, Object]): PipelineStage = {
- val maxIter = TrainParamReader.readParam(params,TrainParamEnum.MAXI_ITER).asInstanceOf[Int]
- val regParam = TrainParamReader.readParam(params,TrainParamEnum.REG_PARAM).asInstanceOf[Double]
- val elasticNetParam = TrainParamReader.readParam(params,TrainParamEnum.ELASTIC_NET_PARAM).asInstanceOf[Double]
- val lr = new LogisticRegression()
- .setLabelCol("label")
- .setFeaturesCol("indexedFeatures")
- .setMaxIter(maxIter)
- .setRegParam(regParam)
- .setElasticNetParam(elasticNetParam)
- lr
- }
- }
|