Przeglądaj źródła

feat:优化了模型的自动存储

DingXiaoYu 3 lat temu
rodzic
commit
ea996d0ae8

+ 3 - 5
java_compile/src/main/java/com/example/services/ModelService.java

@@ -266,11 +266,10 @@ public class ModelService {
                             Float.parseFloat(modelArguments.get("RegParam")),
                             Float.parseFloat(modelArguments.get("ElasticNetParam")));
                     log.info("\n\nstore Model result\n\n");
-                    HashMap<String,String> lrRst = lrClassification.lrResult(lrModel);
+                    HashMap<String,String> lrRst = lrClassification.lrResult(lrModel,modelId);
 
                     log.info("\n\nupdate result of model\n\n");
                     rsl = updateResOfModel(modelId, lrRst, lrModel);
-                    utils.storeLogisticModel(lrModel,modelId);
                     break;
                 case "DecisionTree":
             /*
@@ -298,7 +297,6 @@ public class ModelService {
                     HashMap<String,String> rfResult = rfClassification.dtResult(rfModel, modelId);
 
                     rsl = updateResOfModel(modelId,rfResult,rfModel);
-                    utils.storeModel(rfModel,modelId);
                     break;
                 case "GBDT":
                     /*
@@ -319,7 +317,7 @@ public class ModelService {
                     KMeansModel kmeansModel = kmeansCluster.dtTraining(libsvm,
                             Integer.parseInt(modelArguments.get("NumOfCluster")),
                             Integer.parseInt(modelArguments.get("Seed")));
-                    HashMap<String,String> kmeansResult = kmeansCluster.dtResult(kmeansModel);
+                    HashMap<String,String> kmeansResult = kmeansCluster.dtResult(kmeansModel,modelId);
                     rsl = updateResOfModel(modelId,kmeansResult,kmeansModel);
                     break;
 
@@ -344,7 +342,7 @@ public class ModelService {
                             Integer.parseInt(modelArguments.get("Category")),
                             Float.parseFloat(modelArguments.get("TrainingDataSetOccupy"))
                             );
-                    HashMap<String,String> rfRegResult = rfRegression.dtResult(rfReg);
+                    HashMap<String,String> rfRegResult = rfRegression.dtResult(rfReg,modelId);
 
                     rsl = updateResOfModel(modelId,rfRegResult,rfReg);
                     break;

+ 1 - 1
java_compile/src/main/scala/com/example/model/classification/GDBTClassification.scala

@@ -121,7 +121,7 @@ class GDBTClassification extends SparkConnect{
     rslMap.put("f1", f1.toString)
 
     if(utils.storeModel(pipelineModel, modelId))
-      rslMap
+      return rslMap
     return null
   }
 

+ 10 - 2
java_compile/src/main/scala/com/example/model/classification/LRClassification.scala

@@ -3,7 +3,9 @@ package com.example.model.classification
 import java.util
 
 import com.example.SparkConnect
+import com.example.model.helper.Utils
 import org.apache.spark.ml.classification.{BinaryLogisticRegressionSummary, LogisticRegression, LogisticRegressionModel}
+import org.springframework.beans.factory.annotation.Autowired
 import org.springframework.stereotype.Component
 
 import collection.JavaConverters._
@@ -15,6 +17,9 @@ import collection.JavaConverters._
 @Component
 class LRClassification extends SparkConnect{
 
+  @Autowired
+  val utils: Utils = null
+
   /*
   * initial params
   * */
@@ -53,7 +58,7 @@ class LRClassification extends SparkConnect{
     * @return
     */
   @throws(classOf[Exception])
-  def lrResult(lrModel: LogisticRegressionModel): util.HashMap[String, String] = {
+  def lrResult(lrModel: LogisticRegressionModel, modelId: Int): util.HashMap[String, String] = {
     val rslMap = new util.HashMap[String, String]()
     val trainingSummary = lrModel.summary
 
@@ -73,7 +78,10 @@ class LRClassification extends SparkConnect{
 
     println(binarySummary.areaUnderROC)
     rslMap.put("areaUnderROC", binarySummary.areaUnderROC.toString)
-    rslMap
+
+    if(utils.storeLogisticModel(lrModel, modelId))
+      return rslMap
+    return null
   }
 
 }

+ 11 - 2
java_compile/src/main/scala/com/example/model/cluster/KmeansCluster.scala

@@ -3,11 +3,13 @@ package com.example.model.cluster
 import java.util
 
 import com.example.SparkConnect
+import com.example.model.helper.Utils
 import org.apache.spark.ml.PipelineModel
 import org.apache.spark.sql.{Dataset, Row}
 import org.apache.spark.ml.clustering.KMeansModel
 import org.apache.spark.ml.clustering.KMeans
 import org.apache.spark.ml.feature.{StringIndexer, VectorIndexer}
+import org.springframework.beans.factory.annotation.Autowired
 import org.springframework.stereotype.Component
 import spire.std.float
 
@@ -18,6 +20,10 @@ import spire.std.float
 class KmeansCluster extends SparkConnect{
 
   var trainingSet:Dataset[Row]  = _
+
+  @Autowired
+  val utils: Utils = null
+
   def dtTraining(libsvmFile: String,
                  numK: Int,
                  seed: Int): KMeansModel={
@@ -56,7 +62,7 @@ class KmeansCluster extends SparkConnect{
   }
 
   @throws(classOf[Exception])
-  def dtResult(kmeansModel: KMeansModel): util.HashMap[String, String] = {
+  def dtResult(kmeansModel: KMeansModel, modelId: Int): util.HashMap[String, String] = {
     val resList = new util.HashMap[String, String]()
     val res = kmeansModel.clusterCenters
     res.toString
@@ -67,6 +73,9 @@ class KmeansCluster extends SparkConnect{
       resList.put(num.toString,elem.toArray.mkString(","))
     }
     println(resList)
-    resList
+
+    if(utils.storeKMeansModel(kmeansModel, modelId))
+      return resList
+    return null
   }
 }

+ 22 - 18
java_compile/src/main/scala/com/example/model/helper/Utils.scala

@@ -1,10 +1,10 @@
 package com.example.model.helper
 
-import java.io.{ByteArrayOutputStream, ObjectOutputStream}
 import com.example.SparkConnect
 import com.example.dao.{HdfsDao, ModelDao}
+import org.apache.spark.ml.PipelineModel
 import org.apache.spark.ml.classification.LogisticRegressionModel
-import org.apache.spark.ml.{Model, PipelineModel}
+import org.apache.spark.ml.clustering.KMeansModel
 import org.springframework.beans.factory.annotation.Autowired
 import org.springframework.stereotype.Component
 
@@ -37,31 +37,35 @@ class Utils extends SparkConnect{
     if (modelObject != null) return true
 
     return false
-//    val os = new ByteArrayOutputStream() // 定义一个字节数组输出流
-//    val out = new ObjectOutputStream(os) // 对象输出流
-//    out.writeObject(pipelineModel)
-//    val modelByte = os.toByteArray // byte[]
-//    var modelFile = new ModelFile() // 申明一个模型对象
-//    modelFile.setModel(modelByte)
-//    os.close()
-//    out.close()
-//    val modelFileId = modelFileDao.save(modelFile).getId
-//    val modelOfGet = modelDao.findById(modelId)
-//    modelOfGet.setModelPath(modelFileId.toString)
-//    val modelObject = modelDao.save(modelOfGet)
-//    if (modelObject != null) return true
-//    return false
   }
 
 
   @throws(classOf[Exception])
-  def storeLogisticModel(pipelineModel: LogisticRegressionModel, modelId: Int): Boolean = {
+  def storeLogisticModel(logisticRegressionModel: LogisticRegressionModel, modelId: Int): Boolean = {
     val modelPath = hdfsServer + "_" + modelId + ".model"
 
 
     hdfsDao.deleteFileInHdfs(modelPath, true)
 
-    val flag = pipelineModel.save(modelPath)
+    val flag = logisticRegressionModel.save(modelPath)
+    val modelOfGet = modelDao.findById(modelId)
+
+    modelOfGet.setModelPath(modelPath)
+
+    val modelObject = modelDao.save(modelOfGet)
+    if (modelObject != null) return true
+
+    return false
+  }
+
+  @throws(classOf[Exception])
+  def storeKMeansModel(kMeansModel: KMeansModel, modelId: Int): Boolean = {
+    val modelPath = hdfsServer + "_" + modelId + ".model"
+
+
+    hdfsDao.deleteFileInHdfs(modelPath, true)
+
+    val flag = kMeansModel.save(modelPath)
     val modelOfGet = modelDao.findById(modelId)
 
     modelOfGet.setModelPath(modelPath)

+ 12 - 4
java_compile/src/main/scala/com/example/model/regression/RFRegression.scala

@@ -3,12 +3,13 @@ package com.example.model.regression
 import java.util
 
 import com.example.SparkConnect
+import com.example.model.helper.Utils
 import org.apache.spark.ml.{Pipeline, PipelineModel}
-
-import org.apache.spark.ml.evaluation.{RegressionEvaluator}
+import org.apache.spark.ml.evaluation.RegressionEvaluator
 import org.apache.spark.ml.feature.{IndexToString, StringIndexer, VectorIndexer}
 import org.apache.spark.ml.regression.{RandomForestRegressionModel, RandomForestRegressor}
 import org.apache.spark.sql.{Dataset, Row}
+import org.springframework.beans.factory.annotation.Autowired
 import org.springframework.stereotype.Component
 
 /**
@@ -19,6 +20,10 @@ import org.springframework.stereotype.Component
 class RFRegression extends SparkConnect{
   var trainingSet:Dataset[Row]  = _
   var testData:Dataset[Row] = _
+
+  @Autowired
+  val utils: Utils = null
+
   @throws(classOf[Exception])
   def dtTraining(libsvmFile: String,
                  category: Int,
@@ -72,7 +77,7 @@ class RFRegression extends SparkConnect{
 
 
   @throws(classOf[Exception])
-  def dtResult(pipelineModel: PipelineModel): util.HashMap[String, String] = {
+  def dtResult(pipelineModel: PipelineModel, modelId: Int): util.HashMap[String, String] = {
     val rslMap = new util.HashMap[String, String]()
     // Make predictions.
     val predictions = pipelineModel.transform(testData)
@@ -101,7 +106,10 @@ class RFRegression extends SparkConnect{
 
     rslMap.put("rmse", rmse.toString)
     println(rslMap.get("rmse"))
-    rslMap
+
+    if(utils.storeModel(pipelineModel, modelId))
+      return rslMap
+    return null
   }
 
 }