Explorar el Código

fix: 可以处理训练集未出现的标签,加入了handleinvalid = keep

191250028 hace 3 años
padre
commit
4dc069e56e

+ 1 - 1
java_compile/src/main/java/com/example/util/TrainParamLoader.java

@@ -24,7 +24,7 @@ public class TrainParamLoader {
 
     public static Map<String,Object> getParams(Map<String,String> modelArguments,String modelType){
         Map<String,Object> params = new HashMap<>();
-        TrainParamEnum[] toFill = modelTrainParamMap.get(modelType);
+        TrainParamEnum[] toFill = modelTrainParamMap.getOrDefault(modelType,new TrainParamEnum[]{});
         for(TrainParamEnum trainParamEnum : toFill){
             String paramName = trainParamEnum.getParamName();
             String paramType = trainParamEnum.getParamType();

+ 1 - 0
java_compile/src/main/scala/com/example/data/DataSetHelper.scala

@@ -163,6 +163,7 @@ class DataSetHelper extends SparkConnect{
           new StringIndexer()
             .setInputCol(stringTypeFeature)
             .setOutputCol(stringTypeFeature + "_index")
+            .setHandleInvalid("keep")/* Unseen label: 黑. To handle unseen labels, set Param handleInvalid to keep.*/
             .fit(dataSet)
         })
         pipelineModel = new Pipeline().setStages(stringTypeFeatureIndexers.toArray).fit(dataSet)

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

@@ -19,7 +19,7 @@ class DTClassification extends DefaultClassification {
 
   override def buildClassifier(params: util.Map[String, Object]): PipelineStage = {
     val dt = new DecisionTreeClassifier()
-      .setLabelCol("indexedLabel")
+      .setLabelCol("label")
       .setFeaturesCol("indexedFeatures")
     dt
   }

+ 6 - 2
java_compile/src/main/scala/com/example/model/classification/DefaultClassification.scala

@@ -39,7 +39,7 @@ class DefaultClassification extends SparkConnect {
     val featureIndexer = new VectorIndexer()
       .setInputCol("features")
       .setOutputCol("indexedFeatures")
-      .setMaxCategories(maxCategories)
+      .setMaxCategories(maxCategories).setHandleInvalid("keep")
 
     trainSet.show()
 
@@ -79,7 +79,7 @@ class DefaultClassification extends SparkConnect {
     val f1 = evaluator.setMetricName("f1").evaluate(predictions);
 
 
-    val model = pipelineModel.stages(1).asInstanceOf[LogisticRegressionModel]
+    val model = pipelineModel.stages(1).asInstanceOf[classifier.type]
     //println("Learned classification tree model:\n" + treeModel.toDebugString)
     log.info("\n\nLearned classification Logistic model:\n" + model.toString());
 
@@ -103,5 +103,9 @@ class DefaultClassification extends SparkConnect {
     null
   }
 
+  def getClassifierModel(): Unit ={
+
+  }
+
 
 }

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

@@ -11,7 +11,7 @@ import org.springframework.stereotype.Component
 class GDBTClassification extends DefaultClassification{
   override def buildClassifier(params: util.Map[String, Object]): PipelineStage = {
     // Train a gdbt model.
-    val gbt = new GBTClassifier().setLabelCol("indexedLabel")
+    val gbt = new GBTClassifier().setLabelCol("label")
       .setFeaturesCol("indexedFeatures")
     gbt
   }

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

@@ -24,7 +24,7 @@ class MulPerClassification extends DefaultClassification{
     val layers = Array[Int](2, 5, 4, 5)
     layers(0) = featureNum
     val mlpc = new MultilayerPerceptronClassifier()
-      .setLabelCol("indexedLabel")
+      .setLabelCol("label")
       .setFeaturesCol("indexedFeatures")
       .setLayers(layers)
       .setBlockSize(128)

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

@@ -19,7 +19,7 @@ import com.example.util.Constant
 class NaiveBayesClassification extends DefaultClassification{
 
   override def buildClassifier(params: util.Map[String, Object]): PipelineStage = {
-    val nbc = new NaiveBayes().setLabelCol("indexedLabel")
+    val nbc = new NaiveBayes().setLabelCol("label")
       .setFeaturesCol("indexedFeatures")
     nbc
   }

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

@@ -20,7 +20,7 @@ class RFClassification extends  DefaultClassification{
 
   override def buildClassifier(params: util.Map[String, Object]): PipelineStage = {
     val rf = new RandomForestClassifier()
-      .setLabelCol("indexedLabel")
+      .setLabelCol("label")
       .setFeaturesCol("indexedFeatures")
 
     rf

+ 13 - 0
java_compile/src/main/scala/com/example/util/TrainParamReader.scala

@@ -2,8 +2,21 @@ package com.example.util
 
 import java.util
 
+/**
+ * 可以直接使用,相当于静态方法
+ * @author   fanyanpeng
+ * @date 2023/3/22 21:24
+ */
 object TrainParamReader {
 
+  /**
+   * 获取具体的参数,并进行缺失值补充
+   * @author   fanyanpeng
+   * @date 2023/3/22 21:24
+   * @param params 参数表
+   * @param trainParamEnum 训练参数枚举类
+   * @return java.lang.Object
+   */
   def readParam(params: util.Map[String, Object],trainParamEnum: TrainParamEnum): Any = {
     return params.getOrDefault(trainParamEnum.getParamName, trainParamEnum.getDefaultValue)
   }