Browse Source

feat: 添加Logistic模型的存储

191250028 3 years ago
parent
commit
a3c8f1bcf3

+ 6 - 0
java_compile/a.csv

@@ -0,0 +1,6 @@
+ score,label
+1,1
+2,1
+3,0
+4,0
+5,0

+ 1 - 0
java_compile/src/main/java/com/example/services/ModelService.java

@@ -270,6 +270,7 @@ public class ModelService {
 
                     log.info("\n\nupdate result of model\n\n");
                     rsl = updateResOfModel(modelId, lrRst, lrModel);
+                    utils.storeLogisticModel(lrModel,modelId);
                     break;
                 case "DecisionTree":
             /*

+ 21 - 2
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.springframework.beans.factory.annotation.Autowired
 import org.springframework.stereotype.Component
 
@@ -52,4 +52,23 @@ class Utils extends SparkConnect{
 //    if (modelObject != null) return true
 //    return false
   }
+
+
+  @throws(classOf[Exception])
+  def storeLogisticModel(pipelineModel: LogisticRegressionModel, modelId: Int): Boolean = {
+    val modelPath = hdfsServer + "_" + modelId + ".model"
+
+
+    hdfsDao.deleteFileInHdfs(modelPath, true)
+
+    val flag = pipelineModel.save(modelPath)
+    val modelOfGet = modelDao.findById(modelId)
+
+    modelOfGet.setModelPath(modelPath)
+
+    val modelObject = modelDao.save(modelOfGet)
+    if (modelObject != null) return true
+
+    return false
+  }
 }