Преглед изворни кода

feat: 自动生成数据集合支持中文(需删除数据库,重新执行spring)
feat: 修改模型上传时机

191250028 пре 3 година
родитељ
комит
085a5d5e69

+ 15 - 0
java_compile/src/main/java/com/example/config/MysqlDialog.java

@@ -0,0 +1,15 @@
+package com.example.config;
+
+import org.hibernate.dialect.MySQL5Dialect;
+
+/**
+ * @author fanyanpeng
+ * @date 2023/2/28 18:54
+ */
+public class MysqlDialog extends MySQL5Dialect {
+
+    @Override
+    public String getTableTypeString() {
+        return " ENGINE=InnoDB DEFAULT CHARSET=utf8mb4";
+    }
+}

+ 6 - 2
java_compile/src/main/java/com/example/services/ModelService.java

@@ -218,7 +218,11 @@ public class ModelService {
      * @param model
      * @return
      */
-    public boolean updateResOfModel(int modelId, HashMap<String, String> resOfModel, Object model){
+    public boolean updateResOfModel(int modelId, HashMap<String, String> resOfModel, Object model) throws Exception {
+
+        // 将模型持久化(上传到hdfs)
+        utils.storeModel(model,modelId);
+
         Model modelOfGet =  modelDao.findById(modelId);
         //log.info("\n\nmodel:"+JSON.toJSONString(model)+"\n\n");
         //T.B.D.
@@ -375,7 +379,7 @@ public class ModelService {
 
     public List<Map<String, String>> useModel(int fileId, int modelId) throws Exception{
         List<Map<String,String>> res = fitTestData.transformData(fileId,modelId);
-        fileService.deleteFile(fileId);
+//        fileService.deleteFile(fileId);
         return res;
     }
 

+ 4 - 1
java_compile/src/main/resources/application.properties

@@ -10,7 +10,7 @@ com.jtang.spark.serverAddr = local
 
 com.jtang.spark.appName = java_compile
 
-spring.datasource.url=jdbc:mysql://localhost/AI_Program?serverTimezone=Asia/Shanghai
+spring.datasource.url=jdbc:mysql://localhost/AI_Program?characterEncoding=utf8&&serverTimezone=Asia/Shanghai
 spring.datasource.username=root
 spring.datasource.password=123456
 spring.datasource.driver-class-name=com.mysql.cj.jdbc.Driver
@@ -18,6 +18,9 @@ spring.datasource.driver-class-name=com.mysql.cj.jdbc.Driver
 spring.datasource.data=classpath:data.sql
 spring.datasource.initialization-mode=always
 
+spring.jpa.properties.hibernate.dialect=com.example.config.MysqlDialog
+
+
 spring.jpa.hibernate.ddl-auto=update
 spring.jpa.show-sql=true
 

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

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

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

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

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

@@ -79,9 +79,7 @@ class LRClassification extends SparkConnect{
     println(binarySummary.areaUnderROC)
     rslMap.put("areaUnderROC", binarySummary.areaUnderROC.toString)
 
-    if(utils.storeLogisticModel(lrModel, modelId))
-      return rslMap
-    return null
+    return rslMap
   }
 
 }

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

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

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

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

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

@@ -108,8 +108,6 @@ class RFClassification extends SparkConnect{
     rslMap.put("weightedPrecision", weightedPrecision.toString)
     rslMap.put("weightedRecall", weightedRecall.toString)
     rslMap.put("f1", f1.toString)
-    if(utils.storeModel(pipelineModel, modelId))
-      return rslMap
-    return null
+    return rslMap
   }
 }

+ 4 - 6
java_compile/src/main/scala/com/example/model/cluster/KmeansCluster.scala

@@ -63,19 +63,17 @@ class KmeansCluster extends SparkConnect{
 
   @throws(classOf[Exception])
   def dtResult(kmeansModel: KMeansModel, modelId: Int): util.HashMap[String, String] = {
-    val resList = new util.HashMap[String, String]()
+    val rslMap = new util.HashMap[String, String]()
     val res = kmeansModel.clusterCenters
     res.toString
     var num = 0
     for (elem <- res) {
       num =num+1;
 
-      resList.put(num.toString,elem.toArray.mkString(","))
+      rslMap.put(num.toString,elem.toArray.mkString(","))
     }
-    println(resList)
+    println(rslMap)
 
-    if(utils.storeKMeansModel(kmeansModel, modelId))
-      return resList
-    return null
+    return rslMap
   }
 }

+ 29 - 39
java_compile/src/main/scala/com/example/model/helper/Utils.scala

@@ -22,34 +22,41 @@ class Utils extends SparkConnect{
 
 
   @throws(classOf[Exception])
-  def storeModel(pipelineModel: PipelineModel, modelId: Int): Boolean = {
+  def storeModel(anyModel: Any, 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
-  }
-
-
-  @throws(classOf[Exception])
-  def storeLogisticModel(logisticRegressionModel: LogisticRegressionModel, modelId: Int): Boolean = {
-    val modelPath = hdfsServer + "_" + modelId + ".model"
-
-
-    hdfsDao.deleteFileInHdfs(modelPath, true)
+    val modelClassName = anyModel.getClass.getSimpleName
+    val modelClass = anyModel.getClass
+
+    log.info("模型类名:"+modelClassName)
+    log.info("模型类:"+modelClass)
+
+    var flag:Any = null
+
+    modelClassName match {
+      case "PipelineModel"=> {
+        val model = anyModel.asInstanceOf[PipelineModel]
+        flag = model.save(modelPath)
+        log.info("导出PipelineModel, flag="+flag.toString)
+      }
+      case "KMeansModel" => {
+        val model = anyModel.asInstanceOf[KMeansModel]
+        flag = model.save(modelPath)
+        log.info("导出KMeansModel, flag="+flag.toString)
+      }
+      case "LogisticRegressionModel" => {
+        val model = anyModel.asInstanceOf[LogisticRegressionModel]
+        flag = model.save(modelPath)
+        log.info("导出LogisticRegressionModel, flag="+flag.toString)
+      }
+      case _ =>{
+        log.error("没有匹配的model类")
+      }
+    }
 
-    val flag = logisticRegressionModel.save(modelPath)
     val modelOfGet = modelDao.findById(modelId)
-
     modelOfGet.setModelPath(modelPath)
 
     val modelObject = modelDao.save(modelOfGet)
@@ -58,21 +65,4 @@ class Utils extends SparkConnect{
     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)
-
-    val modelObject = modelDao.save(modelOfGet)
-    if (modelObject != null) return true
-
-    return false
-  }
 }

+ 1 - 3
java_compile/src/main/scala/com/example/model/regression/RFRegression.scala

@@ -107,9 +107,7 @@ class RFRegression extends SparkConnect{
     rslMap.put("rmse", rmse.toString)
     println(rslMap.get("rmse"))
 
-    if(utils.storeModel(pipelineModel, modelId))
-      return rslMap
-    return null
+    return rslMap
   }
 
 }