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 } }