|
|
@@ -16,20 +16,16 @@ import com.example.model.regression.RFRegression;
|
|
|
import com.example.services.pojo.DLGeneralModelPojo;
|
|
|
import com.example.services.pojo.DLModelPojo;
|
|
|
import com.example.services.pojo.ModelPojo;
|
|
|
-import com.example.util.Constant;
|
|
|
import com.example.util.Json2Object;
|
|
|
+import com.example.util.TrainParamLoader;
|
|
|
import lombok.Data;
|
|
|
import lombok.extern.slf4j.Slf4j;
|
|
|
import org.apache.spark.ml.PipelineModel;
|
|
|
-import org.apache.spark.ml.classification.LogisticRegressionModel;
|
|
|
-import org.apache.spark.ml.clustering.KMeansModel;
|
|
|
import org.springframework.beans.factory.annotation.Autowired;
|
|
|
import org.springframework.stereotype.Component;
|
|
|
import scala.Tuple2;
|
|
|
|
|
|
-import javax.persistence.Tuple;
|
|
|
import java.util.ArrayList;
|
|
|
-import java.util.HashMap;
|
|
|
import java.util.List;
|
|
|
import java.util.Map;
|
|
|
import java.util.stream.Collectors;
|
|
|
@@ -267,108 +263,16 @@ public class ModelService {
|
|
|
|
|
|
|
|
|
String modelDetailName = modelType.getModelDetailName();
|
|
|
- HashMap<String,String> modelArguments = new Json2Object().Json2HashMap(trainingModel.getArguments());
|
|
|
+ Map<String,String> modelArguments = new Json2Object().Json2HashMap(trainingModel.getArguments());
|
|
|
boolean rsl = false;
|
|
|
+ Map<String,Object> params = TrainParamLoader.getParams(modelArguments,modelDetailName);
|
|
|
try {
|
|
|
- switch (modelDetailName) {
|
|
|
- case "LogisticRegression":
|
|
|
- /*
|
|
|
- * maxIter: Int,
|
|
|
- regParam: Float,
|
|
|
- elasticNetParam: Float
|
|
|
- * */
|
|
|
-
|
|
|
- Map<String,Object> params = new HashMap<>();
|
|
|
- params.put(Constant.TRAIN_PARAM_MAX_CATEGORY,4);
|
|
|
- params.put(Constant.TRAIN_PARAM_MAX_ITER,Integer.parseInt(modelArguments.get("MaxIter")));
|
|
|
- params.put(Constant.TRAIN_PARAM_REG_PARAM,Float.parseFloat(modelArguments.get("RegParam")));
|
|
|
- params.put(Constant.TRAIN_PARAM_ELASTICNET_PARAM,Float.parseFloat(modelArguments.get("ElasticNetParam")));
|
|
|
- Tuple2<PipelineModel,Map<String,String>> trainResult =lrClassification.training(libsvm,params);
|
|
|
-
|
|
|
- PipelineModel lrModel = trainResult._1;
|
|
|
- Map<String,String> lrRst = trainResult._2;
|
|
|
-
|
|
|
- log.info("\n\nupdate result of model\n\n");
|
|
|
- rsl = updateResOfModel(modelId, lrRst, lrModel);
|
|
|
- break;
|
|
|
- case "DecisionTree":
|
|
|
- /*
|
|
|
- * maxIter: Int,
|
|
|
- regParam: Float,
|
|
|
- elasticNetParam: Float
|
|
|
- * */
|
|
|
- PipelineModel dtModel = dtClassification.dtTraining(libsvm,
|
|
|
- Integer.parseInt(modelArguments.get("Category")),
|
|
|
- Float.parseFloat(modelArguments.get("TrainingDataSetOccupy")));
|
|
|
- HashMap<String,String> dtResult = dtClassification.dtResult(dtModel, modelId);
|
|
|
-
|
|
|
- rsl = updateResOfModel(modelId, dtResult, dtModel);
|
|
|
- break;
|
|
|
- case "RandomForest":
|
|
|
- /*
|
|
|
- * maxIter: Int,
|
|
|
- regParam: Float,
|
|
|
- elasticNetParam: Float
|
|
|
- * */
|
|
|
-
|
|
|
- PipelineModel rfModel = rfClassification.dtTraining(libsvm,
|
|
|
- Integer.parseInt(modelArguments.get("Category")),
|
|
|
- Float.parseFloat(modelArguments.get("TrainingDataSetOccupy")));
|
|
|
- HashMap<String,String> rfResult = rfClassification.dtResult(rfModel, modelId);
|
|
|
-
|
|
|
- rsl = updateResOfModel(modelId,rfResult,rfModel);
|
|
|
- break;
|
|
|
- case "GBDT":
|
|
|
- /*
|
|
|
- * maxIter: Int,
|
|
|
- regParam: Float,
|
|
|
- elasticNetParam: Float
|
|
|
- * */
|
|
|
-
|
|
|
- PipelineModel gbtModel = gdbtClassification.dtTraining(libsvm,
|
|
|
- Integer.parseInt(modelArguments.get("Category")),
|
|
|
- Float.parseFloat(modelArguments.get("TrainingDataSetOccupy")));
|
|
|
- HashMap<String,String> gbtResult = gdbtClassification.dtResult(gbtModel, modelId);
|
|
|
-
|
|
|
- rsl = updateResOfModel(modelId,gbtResult,gbtModel);
|
|
|
- break;
|
|
|
-
|
|
|
- case "K-Means":
|
|
|
- KMeansModel kmeansModel = kmeansCluster.dtTraining(libsvm,
|
|
|
- Integer.parseInt(modelArguments.get("NumOfCluster")),
|
|
|
- Integer.parseInt(modelArguments.get("Seed")));
|
|
|
- HashMap<String,String> kmeansResult = kmeansCluster.dtResult(kmeansModel,modelId);
|
|
|
- rsl = updateResOfModel(modelId,kmeansResult,kmeansModel);
|
|
|
- break;
|
|
|
-
|
|
|
- case "MultilayerPerceptronClassifier":
|
|
|
- PipelineModel mlpc = mulPerClassification.dtTraining(libsvm,
|
|
|
- Integer.parseInt(modelArguments.get("Category")),
|
|
|
- Float.parseFloat(modelArguments.get("TrainingDataSetOccupy")), featureNum);
|
|
|
- HashMap<String,String> mlpcResult = mulPerClassification.dtResult(mlpc, modelId);
|
|
|
- rsl = updateResOfModel(modelId,mlpcResult,mlpc);
|
|
|
- break;
|
|
|
-
|
|
|
- case "NaiveBayes":
|
|
|
- PipelineModel nby = naiveBayesClassification.dtTraining(libsvm,
|
|
|
- Integer.parseInt(modelArguments.get("Category")),
|
|
|
- Float.parseFloat(modelArguments.get("TrainingDataSetOccupy")));
|
|
|
- HashMap<String,String> nbyResult = naiveBayesClassification.dtResult(nby, modelId);
|
|
|
- rsl = updateResOfModel(modelId,nbyResult,nby);
|
|
|
- break;
|
|
|
-
|
|
|
- case "RandomForestRegression":
|
|
|
- PipelineModel rfReg = rfRegression.dtTraining(libsvm,
|
|
|
- Integer.parseInt(modelArguments.get("Category")),
|
|
|
- Float.parseFloat(modelArguments.get("TrainingDataSetOccupy"))
|
|
|
- );
|
|
|
- HashMap<String,String> rfRegResult = rfRegression.dtResult(rfReg,modelId);
|
|
|
-
|
|
|
- rsl = updateResOfModel(modelId,rfRegResult,rfReg);
|
|
|
- break;
|
|
|
- default:
|
|
|
- rsl = false;
|
|
|
- }
|
|
|
+ Tuple2<PipelineModel,Map<String,String>> trainResult = trainClassifier(libsvm,params,modelDetailName);
|
|
|
+ PipelineModel lrModel = trainResult._1;
|
|
|
+ Map<String,String> lrRst = trainResult._2;
|
|
|
+
|
|
|
+ log.info("\n\nupdate result of model\n\n");
|
|
|
+ rsl = updateResOfModel(modelId, lrRst, lrModel);
|
|
|
} catch (Exception e) {
|
|
|
e.printStackTrace();
|
|
|
changeStatusOfTrained(modelId, false);
|
|
|
@@ -400,4 +304,39 @@ public class ModelService {
|
|
|
}
|
|
|
|
|
|
|
|
|
+ /**
|
|
|
+ * 统一的算法访问接口,仅仅包括分类算法
|
|
|
+ * @author fanyanpeng
|
|
|
+ * @date 2023/3/22 21:15
|
|
|
+ * @param libsvmUrl
|
|
|
+ * @param params
|
|
|
+ * @param modelType
|
|
|
+ * @return scala.Tuple2<org.apache.spark.ml.PipelineModel,java.util.Map<java.lang.String,java.lang.String>>
|
|
|
+ */
|
|
|
+ private Tuple2<PipelineModel,Map<String,String>> trainClassifier(String libsvmUrl, Map<String,Object> params, String modelType) throws Exception {
|
|
|
+ Tuple2<PipelineModel,Map<String,String>> result=null;
|
|
|
+ DefaultClassification defaultClassification=null;
|
|
|
+ switch (modelType){
|
|
|
+ case "LogisticRegression":
|
|
|
+ defaultClassification = lrClassification;break;
|
|
|
+ case "DecisionTree":
|
|
|
+ defaultClassification = dtClassification;break;
|
|
|
+ case "RandomForest":
|
|
|
+ defaultClassification = rfClassification;break;
|
|
|
+ case "GBDT":
|
|
|
+ defaultClassification = gdbtClassification;break;
|
|
|
+ case "MultilayerPerceptronClassifier":
|
|
|
+ defaultClassification = mulPerClassification;break;
|
|
|
+ case "NaiveBayes":
|
|
|
+ defaultClassification = naiveBayesClassification;break;
|
|
|
+ default:
|
|
|
+ break;
|
|
|
+
|
|
|
+ }
|
|
|
+ return defaultClassification.training(libsvmUrl,params);
|
|
|
+ }
|
|
|
+
|
|
|
+
|
|
|
+
|
|
|
+
|
|
|
}
|