Bläddra i källkod

feat:实现了所有模型的使用

DingXiaoYu 3 år sedan
förälder
incheckning
c891c84c78
1 ändrade filer med 51 tillägg och 19 borttagningar
  1. 51 19
      java_compile/src/main/scala/com/example/data/FitTestData.scala

+ 51 - 19
java_compile/src/main/scala/com/example/data/FitTestData.scala

@@ -2,15 +2,16 @@ package com.example.data
 
 import java.util
 
-import com.alibaba.fastjson.{JSON, JSONArray, JSONObject}
 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.{StringIndexer, StringIndexerModel}
 import org.apache.spark.ml.{Pipeline, PipelineModel}
-import org.apache.spark.ml.feature.{StringIndexer, StringIndexerModel, VectorIndexer}
-import org.apache.spark.sql.{DataFrame, SaveMode}
 import org.apache.spark.sql.types.StringType
+import org.apache.spark.sql.{DataFrame, SaveMode}
 import org.springframework.beans.factory.annotation.Autowired
 import org.springframework.stereotype.Component
 
@@ -63,8 +64,7 @@ class FitTestData extends SparkConnect{
     val keyFeature = headerInfoDao.findById(Integer.parseInt(fieldsId.get(fieldsId.size()-1))).getFieldName
     var i: Int = 0
     var flag = false
-    while(iter.hasNext && flag==false){
-
+    while(iter.hasNext && !flag){
       fieldNames(i) = headerInfoDao.findById(Integer.parseInt(iter.next())).getFieldName
       println(fieldNames(i))
       i = i+1
@@ -137,20 +137,52 @@ class FitTestData extends SparkConnect{
 
     //modelPath : hdfsServer + "_" + modelId + ".model"
     val modelPath = model.getModelPath
-    val model1 = PipelineModel.load(modelPath)
-    val prections = model1.transform(data)
-
-    prections.show()
-
-    val results = prections.select("predictedLabel")
-    val columns = keyFeature
-    val res = results.collect().map(row => {
-      row.toSeq.zipWithIndex.map(pair => {
-        (columns, pair._1.toString)
-      }).toMap.asJava
-    }).toList.asJava
-
-    res
+    if (model.getModelTypeId==1) {
+      val model1=LogisticRegressionModel.load(modelPath)
+      val prections = model1.transform(data)
+
+      prections.show()
+
+      val results = prections.select("predictedLabel")
+      val columns = keyFeature
+      val res = results.collect().map(row => {
+        row.toSeq.zipWithIndex.map(pair => {
+          (columns, pair._1.toString)
+        }).toMap.asJava
+      }).toList.asJava
+
+      res
+    }else if (model.getModelTypeId==5){
+      val model1 = KMeansModel.load(modelPath)
+      val prections = model1.transform(data)
+
+      prections.show()
+
+      val results = prections.select("predictedLabel")
+      val columns = keyFeature
+      val res = results.collect().map(row => {
+        row.toSeq.zipWithIndex.map(pair => {
+          (columns, pair._1.toString)
+        }).toMap.asJava
+      }).toList.asJava
+
+      res
+    }else{
+      val model1 = PipelineModel.load(modelPath)
+      val prections = model1.transform(data)
+
+      prections.show()
+
+      val results = prections.select("predictedLabel")
+      val columns = keyFeature
+      val res = results.collect().map(row => {
+        row.toSeq.zipWithIndex.map(pair => {
+          (columns, pair._1.toString)
+        }).toMap.asJava
+      }).toList.asJava
+
+      res
+    }
   }
 
 }