Prechádzať zdrojové kódy

feat: 可以走完一个流程

191250028 3 rokov pred
rodič
commit
a2fb2c2aa5

+ 5 - 0
java_compile/src/main/scala/com/example/data/DataSetHelper.scala

@@ -4,6 +4,7 @@ import com.example.SparkConnect
 import com.example.dao.{FileInfoDao, HdfsDao}
 import org.apache.spark.ml.{Pipeline, PipelineModel}
 import org.apache.spark.ml.feature.{StringIndexer, VectorAssembler}
+import org.apache.spark.sql.functions.{col, lit}
 import org.apache.spark.sql.types.StringType
 import org.apache.spark.sql.{Column, DataFrame, SaveMode, SparkSession}
 import org.springframework.beans.factory.annotation.Autowired
@@ -206,6 +207,10 @@ class DataSetHelper extends SparkConnect{
       // 单个向量
       dataSet = assembler.transform(dataSet)
         .select("features")
+      dataSet = dataSet.withColumn("label",lit(0.0))
+      dataSet = dataSet.select(col("label"),col("features"))
+
+
     }
 
 

+ 9 - 28
java_compile/src/main/scala/com/example/data/FitTestData.scala

@@ -7,7 +7,7 @@ import com.example.entity.{Config, Model}
 import com.example.util.Json2Object
 import org.apache.spark.ml.classification.LogisticRegressionModel
 import org.apache.spark.ml.clustering.KMeansModel
-import org.apache.spark.ml.feature.{StringIndexer, StringIndexerModel}
+import org.apache.spark.ml.feature.{IndexToString, StringIndexer, StringIndexerModel}
 import org.apache.spark.ml.{Pipeline, PipelineModel}
 import org.apache.spark.sql.functions.lit
 import org.apache.spark.sql.types.StringType
@@ -39,6 +39,8 @@ class FitTestData extends SparkConnect{
   @Autowired
   val headerInfoDao: HeaderInfoDao = null
 
+  @Autowired
+  val dataSetHelper : DataSetHelper = null
 
   @throws(classOf[Exception])
   def transformData(fileId: Int, modelId: Int): util.List[util.Map[String, String]] = {
@@ -80,32 +82,8 @@ class FitTestData extends SparkConnect{
     import spark.implicits._
     val dataFrameFromParquet = sparkSession.read.parquet(filepath)
 
-
-
-    //此处需要做一个适配,若文件此时没有label列,就先加上
-    var dataFrameForPredict:DataFrame = dataFrameFromParquet
-    if(!dataFrameFromParquet.columns.toList.contains(keyFeature)){
-
-      val labelHeaderInfo = headerInfoDao.findById(Integer.parseInt(labelHeaderInfoId))
-      val valueInfo = labelHeaderInfo.getValueInfo
-      val fieldType = labelHeaderInfo.getFieldType
-      val valueArray = json2Object.Json2List(valueInfo)
-      val labelDefaultValueStr = valueArray.get(0)
-
-      val labelDefaultValue = fieldType match {
-        case "string" => labelDefaultValueStr
-        case "int" => Integer.parseInt(labelDefaultValueStr)
-        case "double"  => labelDefaultValueStr.toDouble
-        case _ => log.error("类型出错")
-      }
-
-      dataFrameForPredict = dataFrameFromParquet.withColumn(keyFeature,lit(labelDefaultValue))
-    }
-    //添加部分到此结束
-
-
     //libsvm文件统一资源定位符。
-    val libsvmFileUri = libsvmAdapter.dataFrame2Libsvm(dataFrameForPredict,fileId,keyFeature,fieldNames)
+    val libsvmFileUri = dataSetHelper.processTestDataSet(modelId,fileId,fieldNames)
 
     // Load the data stored in LIBSVM format as a DataFrame.
     val data = sparkSession.read.format("libsvm").load(libsvmFileUri)
@@ -121,7 +99,7 @@ class FitTestData extends SparkConnect{
 
       prections.show()
 
-      val results = prections.select("predictedLabel")
+      val results = prections.select("prediction")
       val columns = keyFeature
       val res = results.collect().map(row => {
         row.toSeq.zipWithIndex.map(pair => {
@@ -136,7 +114,10 @@ class FitTestData extends SparkConnect{
 
       prections.show()
 
-      val results = prections.select("predictedLabel")
+      val results = prections.select("prediction")
+
+
+
       val columns = keyFeature
       val res = results.collect().map(row => {
         row.toSeq.zipWithIndex.map(pair => {

+ 5 - 16
java_compile/src/main/scala/com/example/model/classification/LRClassification.scala

@@ -48,13 +48,6 @@ class LRClassification extends SparkConnect{
     data.show()
 
 
-    // Index labels, adding metadata to the label column.
-    // Fit on whole dataset to include all labels in index.
-    val labelIndexer = new StringIndexer()
-      .setInputCol("label")
-      .setOutputCol("indexedLabel")
-      .fit(data)
-
     // Automatically identify categorical features, and index them.
     val featureIndexer = new VectorIndexer()
       .setInputCol("features")
@@ -69,21 +62,17 @@ class LRClassification extends SparkConnect{
     testData = testSet
 
     val lr = new LogisticRegression()
-      .setLabelCol("indexedLabel")
+      .setLabelCol("label")
       .setFeaturesCol("indexedFeatures")
       .setMaxIter(maxIter)
       .setRegParam(regParam)
       .setElasticNetParam(elasticNetParam)
 
-    // Convert indexed labels back to original labels.
-    val labelConverter = new IndexToString()
-      .setInputCol("prediction")
-      .setOutputCol("predictedLabel")
-      .setLabels(labelIndexer.labels)
+
 
     // Chain indexers and tree in a Pipeline.
     val pipeline = new Pipeline()
-      .setStages(Array(labelIndexer, featureIndexer, lr, labelConverter))
+      .setStages(Array(featureIndexer, lr))
 
     // Train model. This also runs the indexers.
     val model = pipeline.fit(trainingSet)
@@ -107,7 +96,7 @@ class LRClassification extends SparkConnect{
 
     // Select (prediction, true label) and compute test error.
     val evaluator = new MulticlassClassificationEvaluator()
-      .setLabelCol("indexedLabel")
+      .setLabelCol("label")
       .setPredictionCol("prediction")
 
     val weightedRecall = evaluator.setMetricName("weightedRecall").evaluate(predictions);
@@ -119,7 +108,7 @@ class LRClassification extends SparkConnect{
     val f1 = evaluator.setMetricName("f1").evaluate(predictions);
 
 
-    val model = pipelineModel.stages(2).asInstanceOf[LogisticRegressionModel]
+    val model = pipelineModel.stages(1).asInstanceOf[LogisticRegressionModel]
     //println("Learned classification tree model:\n" + treeModel.toDebugString)
     log.info("\n\nLearned classification Logistic model:\n" + model.toString());