Kaynağa Gözat

feat: 合并对训练集和测试集的处理到一份代码

191250028 3 yıl önce
ebeveyn
işleme
5a0ebfbd95

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

@@ -260,7 +260,7 @@ public class ModelService {
         int featureNum = features.size();
 //        String libsvm = libsvmAdapter.parquet2Libsvm(configInfo.getFileInfoId(), keyFeature, features.toArray(new String[0]));
 
-        String libsvm = dataSetHelper.processTrainDataSet(modelId,configInfo.getFileInfoId(),keyFeature, features.toArray(new String[0]));
+        String libsvm = dataSetHelper.processTrainDataSet(modelId,configInfo.getFileInfoId(),features.toArray(new String[0]),keyFeature );
 
 
         String modelDetailName = modelType.getModelDetailName();

+ 179 - 82
java_compile/src/main/scala/com/example/data/DataSetHelper.scala

@@ -9,9 +9,6 @@ import org.apache.spark.sql.{Column, DataFrame, SaveMode, SparkSession}
 import org.springframework.beans.factory.annotation.Autowired
 import org.springframework.stereotype.Component
 
-import javax.annotation.PostConstruct
-
-
 @Component
 class DataSetHelper extends SparkConnect{
 
@@ -21,146 +18,246 @@ class DataSetHelper extends SparkConnect{
   @Autowired
   val hdfsDao: HdfsDao = null
 
-
+  /**
+   * 填写null值,目前处理了StringType
+   * @author   fanyanpeng
+   * @date 2023/3/21 16:00
+   * @param dataFrame 通用的数据
+   * @return org.apache.spark.sql.Dataset<org.apache.spark.sql.Row>
+   */
   def fillNull(dataFrame: DataFrame): DataFrame ={
-
     val stringTypeFeatures = dataFrame.schema.filter(_.dataType.equals(StringType)).map(_.name)
     return dataFrame.na.fill("NA",stringTypeFeatures)
   }
 
+  /**
+   * 获取转换特征pipeline的uri,目前实现在hdfs
+   * @author   fanyanpeng
+   * @date 2023/3/21 16:01
+   * @param modelId 模型id
+   * @return java.lang.String
+   */
   def getStringTypeFeatureTransformerUri(modelId:Int): String ={
     hdfsUri+"pipeline_StringTypeFeatureTransformer_modelId-"+modelId
   }
 
+  /**
+   * 获取文件存储的libsvm的uri,目前实现在hdfs
+   * @author   fanyanpeng
+   * @date 2023/3/21 16:03
+   * @param fileFlag
+   * @return java.lang.String
+   */
   def getFileLibsvmUri(fileFlag: Int): String = {
     hdfsUri + fileFlag + "-2023.libsvm"
   }
 
 
+  /**
+   * 获得标签转换器的uri,存储在hdfs上
+   * @author   fanyanpeng
+   * @date 2023/3/21 16:04
+   * @param modelId
+   * @return java.lang.String
+   */
   def getLabelStringIndexerUri(modelId: Int):String = {
     hdfsUri+"LabelStringIndexer_modelId-"+modelId
   }
 
 
-  @throws(classOf[Exception])
-  def processTrainDataSet(modelId:Int,
-                          fileId: Int,
-                          label: String,
-                          features: Array[String]): String ={
+
+
+
+
+
+
+  /**
+   * 训练数据与测试数据中有大量重复部分,流程一致
+   * 为保证处理的一致性,且充分复用代码,使用主流水线,分支流水线的方式解决
+   * @author   fanyanpeng
+   * @date 2023/3/21 16:05
+   * @param modelId 模型id
+   * @param fileId 文件id
+   * @param features 特征名数组
+   * @param label 标签名
+   * @return java.lang.String
+   */
+  private def processDataSet(modelId: Int,
+                             fileId: Int,
+                             features: Array[String],
+                             label: String = null): String = {
+//    当前数据是否为训练数据
+    val isTrainData = label != null
 
     val spark = sparkSession
     val trainDataFile = fileLocDao.findById(fileId)
     val trainDataFileUri = trainDataFile.getLocation
-    log.info("trainDataFileUri =  "+trainDataFileUri)
-
-    //训练集读取
-    var trainSet:DataFrame = spark.read.parquet(trainDataFileUri)
-
-    //保留标签|特征
-    trainSet = trainSet.select(label,features:_*).persist()
+    log.info("trainDataFileUri =  " + trainDataFileUri)
+
+    var dataSet: DataFrame = spark.read.parquet(trainDataFileUri)
+    log.info("文件加载")
+    dataSet.show(10)
+
+
+//    选取指定列
+    if(isTrainData){
+      // 保留标签
+      // 标签|特征
+      dataSet = dataSet.select(label, features: _*).persist()
+    } else {
+      // 保留特征
+      // 特征
+      dataSet = dataSet.toDF(features: _*).persist()
+    }
+    log.info("剔除其他列")
+    dataSet.show(10)
 
-    //去除标签列为空的数据(无法进行学习)
-    trainSet = trainSet.filter("'"+ label+" ' is not null")
+//    对于训练集,需要剔除没有标签的数据
+    if(isTrainData){
+      // 去除标签列为空的数据(无法进行学习)
+      dataSet = dataSet.filter("'" + label + " ' is not null")
+    }
 
-    //填充特征空值
-    trainSet = fillNull(trainSet)
+//    填充缺省
+    dataSet = fillNull(dataSet)
     log.info("训练集-已填充空值")
-    trainSet.show(5)
+    dataSet.show(10)
 
-    //取出类型为string的特征
-    val stringTypeFeatures = trainSet.schema
-      .filter(_.dataType.equals(StringType))  //是string类型
-      .filter(!_.name.equals(label))   //且不是label列
+    //类型为string的特征
+    val stringTypeFeatures = dataSet.schema
+      .filter(_.dataType.equals(StringType))
+      .filter(!_.name.equals(label)) //此处不处理label列,对于测试集,本行过滤无效
       .map(_.name)
 
-    //字符类型特征 -> 数值类型特征
-    val stringTypeFeatureIndexers = stringTypeFeatures.map(stringTypeFeature => {
-      new StringIndexer()
-        .setInputCol(stringTypeFeature)
-        .setOutputCol(stringTypeFeature + "_index")
-        .fit(trainSet)
-    })
-
-    if(stringTypeFeatures.nonEmpty){ //若存在字符类型特征
-      log.info("需要进行字符型类型转换")
-      log.info("待转换字符型特征有"+stringTypeFeatures.length+"个:")
-      log.info(stringTypeFeatures.toString())
-
-      log.info("构建pipeline进行转换")
-      val pipeline = new Pipeline().setStages(stringTypeFeatureIndexers.toArray)
-      val pipelineModel : PipelineModel = pipeline.fit(trainSet)
+    //若存在字符类型特征
+    if (stringTypeFeatures.nonEmpty) {
+      log.info("待转换字符型特征有" + stringTypeFeatures.length + "个:"+stringTypeFeatures.toString())
 
       val pipelineUri = getStringTypeFeatureTransformerUri(modelId)
-      log.info("转换模型适配完成,正在保存到服务器,保存地址:"+pipelineUri)
-      pipelineModel.write.overwrite().save(pipelineUri)
-
-      trainSet = pipelineModel.transform(trainSet)
+      var pipelineModel:PipelineModel = null;
+      if(isTrainData){
+        //字符类型特征 -> 数值类型特征
+        val stringTypeFeatureIndexers = stringTypeFeatures.map(stringTypeFeature => {
+          new StringIndexer()
+            .setInputCol(stringTypeFeature)
+            .setOutputCol(stringTypeFeature + "_index")
+            .fit(dataSet)
+        })
+        pipelineModel = new Pipeline().setStages(stringTypeFeatureIndexers.toArray).fit(dataSet)
+        pipelineModel.write.overwrite().save(pipelineUri)
+        log.info("转换模型适配完成,已保存到远端:" + pipelineUri)
+      }
+      else {
+        pipelineModel = PipelineModel.load(pipelineUri)
+        log.info("从远程加载完成")
+      }
+
+      //字符串特征转化为数值类型
+      dataSet = pipelineModel.transform(dataSet)
+      log.info("训练集-字符串特征转化为数值类型:")
+      dataSet.show(10)
     }
 
     //训练集-去除非数值的特征
-    trainSet = trainSet.drop(stringTypeFeatures:_*)
+    dataSet = dataSet.drop(stringTypeFeatures: _*)
 
     log.info("训练集-去除非数值的特征:")
-    trainSet.show(5)
-
-
-    //找到为string类型的标签(0个或者1个)
-    val stringLabel = trainSet.schema.filter(_.name.equals(label)).filter(_.dataType.equals(StringType)).map(_.name)
-    if(stringLabel.nonEmpty){ //需要转换,将label的转换方法持久化
-      val labelStringIndexer = new StringIndexer().setInputCol(label).setOutputCol("label").fit(trainSet)
-      val labelStringIndexerUri = getLabelStringIndexerUri(modelId)
-      labelStringIndexer.write.overwrite().save(labelStringIndexerUri)
-      trainSet = labelStringIndexer.transform(trainSet)
-      log.info("label转换模型已经存储到远端,uri:"+labelStringIndexerUri)
-      trainSet.show(5)
+    dataSet.show(10)
+
+
+    val stringLabel = dataSet.schema.filter(_.name.equals(label)).filter(_.dataType.equals(StringType)).map(_.name)
+    if(isTrainData){
+      //找到为string类型的标签(0个或者1个)
+      if (stringLabel.nonEmpty) { //需要转换,将label的转换方法持久化
+        val labelStringIndexer = new StringIndexer().setInputCol(label).setOutputCol("label").fit(dataSet)
+        val labelStringIndexerUri = getLabelStringIndexerUri(modelId)
+        labelStringIndexer.write.overwrite().save(labelStringIndexerUri)
+        dataSet = labelStringIndexer.transform(dataSet)
+        log.info("label转换模型已经存储到远端,uri:" + labelStringIndexerUri)
+        dataSet.show(5)
+      }
     }
 
-    val columns = trainSet.columns
+
+    val columns = dataSet.columns
 
     //过滤标签,取出特征列表
-    val featureList:Array[String] = columns.filter(columnName => (!columnName.equals(label) && !columnName.equals("label")))
+    val featureList: Array[String] = columns.filter(columnName => (!columnName.equals(label) && !columnName.equals("label")))
 
     //将特征列表合成向量
-    val assembler = new VectorAssembler()
+    var assembler = new VectorAssembler()
       .setInputCols(featureList)
       .setOutputCol("features")
 
-    // 若label是字符串类型,那么将选择新的数值型的label列
-    var selectedLabel= label
-    if(stringLabel.nonEmpty){
-      selectedLabel = "label"
-    }
 
-    val assembledData = assembler.transform(trainSet)
-      .select(selectedLabel,"features")
-      .toDF("label","features") //重命名
 
-    log.info("写入libsvm的数据格式如下:")
-    assembledData.show(5)
+    if(isTrainData){
+      // 若label是字符串类型,那么将选择新的数值型的label列
+      var selectedLabel = label
+      if (stringLabel.nonEmpty) {
+        selectedLabel = "label"
+      }
+      dataSet = assembler.transform(dataSet)
+        .select(selectedLabel, "features")
+        .toDF("label", "features") //重命名
 
-    val libsvmFileUri = getFileLibsvmUri(fileId)
+    }
+    else{
+      // 单个向量
+      dataSet = assembler.transform(dataSet)
+        .select("features")
+    }
 
 
-//    implicit val encoder = org.apache.spark.sql.Encoders.STRING
+    log.info("写入libsvm的数据格式如下:")
+    dataSet.show(5)
 
+    val libsvmFileUri = getFileLibsvmUri(fileId)
 
-    assembledData.write.mode(SaveMode.Overwrite) //没有上面那句声明,无法写入成功
+    dataSet.write.mode(SaveMode.Overwrite) //没有上面那句声明,无法写入成功
       .format("libsvm")
       .save(libsvmFileUri)
-
-    log.info("libsvm写入,uri: "+libsvmFileUri)
-
+    log.info("libsvm写入,uri: " + libsvmFileUri)
     return libsvmFileUri
-
   }
 
-  def processTestDataSet(): String = {
 
+  /**
+   * 处理训练数据,返回libsvm的uri
+   *
+   * @author fanyanpeng
+   * @date 2023/3/21 16:04
+   * @param modelId  模型id
+   * @param fileId   文件id
+   * @param features 特征名数组
+   * @param label    标签名
+   * @return java.lang.String
+   */
+  @throws(classOf[Exception])
+  def processTrainDataSet(modelId: Int,
+                          fileId: Int,
+                          features: Array[String],
+                          label: String): String = {
 
-    "libsvmFileUri"
+    return processDataSet(modelId, fileId, features, label)
   }
 
-
+  /**
+   * 处理测试数据
+   *
+   * @author fanyanpeng
+   * @date 2023/3/21 16:05
+   * @param modelId  模型id
+   * @param fileId   文件id
+   * @param features 特征名数组
+   * @return java.lang.String
+   */
+  def processTestDataSet(modelId: Int,
+                         fileId: Int,
+                         features: Array[String]): String = {
+
+    return processDataSet(modelId, fileId, features)
+  }
 
 
 }