|
|
@@ -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)
|
|
|
+ }
|
|
|
|
|
|
|
|
|
}
|