Jelajahi Sumber

feat: 支持复原label,并在本地打印完整测试集和预测结果

191250028 3 tahun lalu
induk
melakukan
c9e91404dc

+ 27 - 6
java_compile/src/main/scala/com/example/data/DataSetHelper.scala

@@ -3,7 +3,7 @@ package com.example.data
 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.ml.feature.{IndexToString, StringIndexer, StringIndexerModel, 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}
@@ -26,7 +26,7 @@ class DataSetHelper extends SparkConnect{
    * @param dataFrame 通用的数据
    * @return org.apache.spark.sql.Dataset<org.apache.spark.sql.Row>
    */
-  def fillNull(dataFrame: DataFrame): DataFrame ={
+  private def fillNull(dataFrame: DataFrame): DataFrame ={
     val stringTypeFeatures = dataFrame.schema.filter(_.dataType.equals(StringType)).map(_.name)
     return dataFrame.na.fill("NA",stringTypeFeatures)
   }
@@ -38,7 +38,7 @@ class DataSetHelper extends SparkConnect{
    * @param modelId 模型id
    * @return java.lang.String
    */
-  def getStringTypeFeatureTransformerUri(modelId:Int): String ={
+  private def getStringTypeFeatureTransformerUri(modelId:Int): String ={
     hdfsUri+"pipeline_StringTypeFeatureTransformer_modelId-"+modelId
   }
 
@@ -49,7 +49,7 @@ class DataSetHelper extends SparkConnect{
    * @param fileFlag
    * @return java.lang.String
    */
-  def getFileLibsvmUri(fileFlag: Int): String = {
+  private def getFileLibsvmUri(fileFlag: Int): String = {
     hdfsUri + fileFlag + "-2023.libsvm"
   }
 
@@ -61,12 +61,31 @@ class DataSetHelper extends SparkConnect{
    * @param modelId
    * @return java.lang.String
    */
-  def getLabelStringIndexerUri(modelId: Int):String = {
+  private def getLabelStringIndexerUri(modelId: Int):String = {
     hdfsUri+"LabelStringIndexer_modelId-"+modelId
   }
 
 
 
+  def restoreLabel(modelId:Int,indexedLabel:DataFrame):DataFrame = {
+
+    var labelDataFrame:DataFrame = indexedLabel.toDF("labelIndexed");
+    try{
+      val labelStringIndexerUri = getLabelStringIndexerUri(modelId)
+      val stringIndexerModel : StringIndexerModel = StringIndexerModel.load(labelStringIndexerUri)
+      val convert:IndexToString = new IndexToString()
+        .setInputCol("labelIndexed")
+        .setOutputCol(stringIndexerModel.inputCol.name)
+        .setLabels(stringIndexerModel.labelsArray(0))
+      labelDataFrame = convert.transform(labelDataFrame)
+        .select(col(stringIndexerModel.inputCol.name))
+    }catch{
+      case e:Exception=>{
+        log.info("转换失败"+e)
+      }
+    }
+    return  labelDataFrame
+  }
 
 
 
@@ -108,7 +127,9 @@ class DataSetHelper extends SparkConnect{
     } else {
       // 保留特征
       // 特征
-      dataSet = dataSet.toDF(features: _*).persist()
+      val columns = dataSet.columns
+      // 需要写成这种格式
+      dataSet = dataSet.select(features.head, features.tail:_*).persist()
     }
     log.info("剔除其他列")
     dataSet.show(10)

+ 56 - 38
java_compile/src/main/scala/com/example/data/FitTestData.scala

@@ -5,13 +5,15 @@ import com.example.SparkConnect
 import com.example.dao._
 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.{IndexToString, StringIndexer, StringIndexerModel}
 import org.apache.spark.ml.{Pipeline, PipelineModel}
-import org.apache.spark.sql.functions.lit
-import org.apache.spark.sql.types.StringType
-import org.apache.spark.sql.{DataFrame, SaveMode}
+import org.apache.spark.rdd.RDD
+import org.apache.spark.sql.functions.{col, lit, typedLit}
+import org.apache.spark.sql.types.{LongType, StringType, StructField, StructType}
+import org.apache.spark.sql.{DataFrame, Row, SaveMode}
 import org.springframework.beans.factory.annotation.Autowired
 import org.springframework.stereotype.Component
 
@@ -62,11 +64,11 @@ class FitTestData extends SparkConnect{
     val iter = fieldsId.listIterator()
     //println(fieldsId(0))
 
-    //val keyFeature = fieldsId.remove(fieldsId.size() - 1)
+    //val labelName = fieldsId.remove(fieldsId.size() - 1)
     val fieldNames: Array[String] =  new Array[String](fieldsId.size()-1)
 
     val labelHeaderInfoId =  fieldsId.get(fieldsId.size()-1)
-    val keyFeature = headerInfoDao.findById(Integer.parseInt(labelHeaderInfoId)).getFieldName
+    val labelName = headerInfoDao.findById(Integer.parseInt(labelHeaderInfoId)).getFieldName
     var i: Int = 0
     var flag = false
     while(iter.hasNext && !flag){
@@ -76,56 +78,72 @@ class FitTestData extends SparkConnect{
       if(i>=fieldsId.size()-1)
         flag = true
     }
+    //上述代码应该被抽象为一个业务逻辑,给定文件,获取文件的分析数据,需要有一个数据结构支撑!
+
+
+
+
+
 
     val spark = sparkSession
 
     import spark.implicits._
-    val dataFrameFromParquet = sparkSession.read.parquet(filepath)
+    val testSet = sparkSession.read.parquet(filepath)
 
     //libsvm文件统一资源定位符。
     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)
-    // Index labels, adding metadata to the label column.
-    // Fit on whole dataset to include all labels in index.
+    // 默认label 和 特征features
+    val testDataLibSvm = sparkSession.read.format("libsvm").load(libsvmFileUri)
 
-
-    //modelPath : hdfsServer + "_" + modelId + ".model"
+    // 完成预测
     val modelPath = model.getModelPath
-    if (model.getModelTypeId==5){
-      val model1 = KMeansModel.load(modelPath)
-      val prections = model1.transform(data)
+    val trainedModel = PipelineModel.load(modelPath)
+    val prediction = trainedModel.transform(testDataLibSvm)
+
+    prediction.show()
+    var predictedLabel = prediction.select("prediction")
+    predictedLabel = dataSetHelper.restoreLabel(modelId, predictedLabel).toDF("prediction_"+labelName)
+
+    val res: util.List[util.Map[String, String]] = predictedLabel.collect().map(row => {
+      row.toSeq.zipWithIndex.map(pair => {
+        (labelName, pair._1.toString)
+      }).toMap.asJava
+    }).toList.asJava
+
+
+    testSet.show()
+    val testSetWithId = addId(testSet)
+    testSetWithId.show()
+    predictedLabel.show()
+    val predictedLabelWithId = addId(predictedLabel)
+    predictedLabelWithId.show()
+
+    var predictTable:DataFrame = testSetWithId.join(predictedLabelWithId,"id");
+    predictTable.show()
+    predictTable = predictTable.orderBy("id")
+    predictTable.show()
+    res
+  }
 
-      prections.show()
 
-      val results = prections.select("prediction")
-      val columns = keyFeature
-      val res = results.collect().map(row => {
-        row.toSeq.zipWithIndex.map(pair => {
-          (columns, pair._1.toString)
-        }).toMap.asJava
-      }).toList.asJava
+  private def addId(dataFrame:DataFrame):DataFrame={
+    val schema: StructType = dataFrame.schema.add(StructField("id", LongType))
+    // DataFrame转RDD 然后调用 zipWithIndex
+    //zipWithIndex返回的是元组数组:Array((Sunday,0), (Monday,1), ...
+    //每一元组,第一个为原来的Row,第二项是index;
+    val dfRDD: RDD[(Row, Long)] = dataFrame.rdd.zipWithIndex()
 
-      res
-    }else{
-      val model1 = PipelineModel.load(modelPath)
-      val prections = model1.transform(data)
+    //合并数组中的元组,第二项需要转换为Row
+    val rowRDD: RDD[Row] = dfRDD.map(tp => Row.merge(tp._1, Row(tp._2)))
 
-      prections.show()
+    // 将添加了索引的RDD 转化为DataFrame
+    val df2 = sparkSession.createDataFrame(rowRDD, schema)
+    return df2;
+  }
 
-      val results = prections.select("prediction")
 
 
 
-      val columns = keyFeature
-      val res = results.collect().map(row => {
-        row.toSeq.zipWithIndex.map(pair => {
-          (columns, pair._1.toString)
-        }).toMap.asJava
-      }).toList.asJava
-      res
-    }
-  }
 
 }