|
|
@@ -5,13 +5,15 @@ import com.example.SparkConnect
|
|
|
import com.example.dao._
|
|
|
import com.example.entity.{Config, Model}
|
|
|
import com.example.util.Json2Object
|
|
|
+
|
|
|
import org.apache.spark.ml.classification.LogisticRegressionModel
|
|
|
import org.apache.spark.ml.clustering.KMeansModel
|
|
|
import org.apache.spark.ml.feature.{IndexToString, StringIndexer, StringIndexerModel}
|
|
|
import org.apache.spark.ml.{Pipeline, PipelineModel}
|
|
|
-import org.apache.spark.sql.functions.lit
|
|
|
-import org.apache.spark.sql.types.StringType
|
|
|
-import org.apache.spark.sql.{DataFrame, SaveMode}
|
|
|
+import org.apache.spark.rdd.RDD
|
|
|
+import org.apache.spark.sql.functions.{col, lit, typedLit}
|
|
|
+import org.apache.spark.sql.types.{LongType, StringType, StructField, StructType}
|
|
|
+import org.apache.spark.sql.{DataFrame, Row, SaveMode}
|
|
|
import org.springframework.beans.factory.annotation.Autowired
|
|
|
import org.springframework.stereotype.Component
|
|
|
|
|
|
@@ -62,11 +64,11 @@ class FitTestData extends SparkConnect{
|
|
|
val iter = fieldsId.listIterator()
|
|
|
//println(fieldsId(0))
|
|
|
|
|
|
- //val keyFeature = fieldsId.remove(fieldsId.size() - 1)
|
|
|
+ //val labelName = fieldsId.remove(fieldsId.size() - 1)
|
|
|
val fieldNames: Array[String] = new Array[String](fieldsId.size()-1)
|
|
|
|
|
|
val labelHeaderInfoId = fieldsId.get(fieldsId.size()-1)
|
|
|
- val keyFeature = headerInfoDao.findById(Integer.parseInt(labelHeaderInfoId)).getFieldName
|
|
|
+ val labelName = headerInfoDao.findById(Integer.parseInt(labelHeaderInfoId)).getFieldName
|
|
|
var i: Int = 0
|
|
|
var flag = false
|
|
|
while(iter.hasNext && !flag){
|
|
|
@@ -76,56 +78,72 @@ class FitTestData extends SparkConnect{
|
|
|
if(i>=fieldsId.size()-1)
|
|
|
flag = true
|
|
|
}
|
|
|
+ //上述代码应该被抽象为一个业务逻辑,给定文件,获取文件的分析数据,需要有一个数据结构支撑!
|
|
|
+
|
|
|
+
|
|
|
+
|
|
|
+
|
|
|
+
|
|
|
|
|
|
val spark = sparkSession
|
|
|
|
|
|
import spark.implicits._
|
|
|
- val dataFrameFromParquet = sparkSession.read.parquet(filepath)
|
|
|
+ val testSet = sparkSession.read.parquet(filepath)
|
|
|
|
|
|
//libsvm文件统一资源定位符。
|
|
|
val libsvmFileUri = dataSetHelper.processTestDataSet(modelId,fileId,fieldNames)
|
|
|
|
|
|
- // Load the data stored in LIBSVM format as a DataFrame.
|
|
|
- val data = sparkSession.read.format("libsvm").load(libsvmFileUri)
|
|
|
- // Index labels, adding metadata to the label column.
|
|
|
- // Fit on whole dataset to include all labels in index.
|
|
|
+ // 默认label 和 特征features
|
|
|
+ val testDataLibSvm = sparkSession.read.format("libsvm").load(libsvmFileUri)
|
|
|
|
|
|
-
|
|
|
- //modelPath : hdfsServer + "_" + modelId + ".model"
|
|
|
+ // 完成预测
|
|
|
val modelPath = model.getModelPath
|
|
|
- if (model.getModelTypeId==5){
|
|
|
- val model1 = KMeansModel.load(modelPath)
|
|
|
- val prections = model1.transform(data)
|
|
|
+ val trainedModel = PipelineModel.load(modelPath)
|
|
|
+ val prediction = trainedModel.transform(testDataLibSvm)
|
|
|
+
|
|
|
+ prediction.show()
|
|
|
+ var predictedLabel = prediction.select("prediction")
|
|
|
+ predictedLabel = dataSetHelper.restoreLabel(modelId, predictedLabel).toDF("prediction_"+labelName)
|
|
|
+
|
|
|
+ val res: util.List[util.Map[String, String]] = predictedLabel.collect().map(row => {
|
|
|
+ row.toSeq.zipWithIndex.map(pair => {
|
|
|
+ (labelName, pair._1.toString)
|
|
|
+ }).toMap.asJava
|
|
|
+ }).toList.asJava
|
|
|
+
|
|
|
+
|
|
|
+ testSet.show()
|
|
|
+ val testSetWithId = addId(testSet)
|
|
|
+ testSetWithId.show()
|
|
|
+ predictedLabel.show()
|
|
|
+ val predictedLabelWithId = addId(predictedLabel)
|
|
|
+ predictedLabelWithId.show()
|
|
|
+
|
|
|
+ var predictTable:DataFrame = testSetWithId.join(predictedLabelWithId,"id");
|
|
|
+ predictTable.show()
|
|
|
+ predictTable = predictTable.orderBy("id")
|
|
|
+ predictTable.show()
|
|
|
+ res
|
|
|
+ }
|
|
|
|
|
|
- prections.show()
|
|
|
|
|
|
- val results = prections.select("prediction")
|
|
|
- val columns = keyFeature
|
|
|
- val res = results.collect().map(row => {
|
|
|
- row.toSeq.zipWithIndex.map(pair => {
|
|
|
- (columns, pair._1.toString)
|
|
|
- }).toMap.asJava
|
|
|
- }).toList.asJava
|
|
|
+ private def addId(dataFrame:DataFrame):DataFrame={
|
|
|
+ val schema: StructType = dataFrame.schema.add(StructField("id", LongType))
|
|
|
+ // DataFrame转RDD 然后调用 zipWithIndex
|
|
|
+ //zipWithIndex返回的是元组数组:Array((Sunday,0), (Monday,1), ...
|
|
|
+ //每一元组,第一个为原来的Row,第二项是index;
|
|
|
+ val dfRDD: RDD[(Row, Long)] = dataFrame.rdd.zipWithIndex()
|
|
|
|
|
|
- res
|
|
|
- }else{
|
|
|
- val model1 = PipelineModel.load(modelPath)
|
|
|
- val prections = model1.transform(data)
|
|
|
+ //合并数组中的元组,第二项需要转换为Row
|
|
|
+ val rowRDD: RDD[Row] = dfRDD.map(tp => Row.merge(tp._1, Row(tp._2)))
|
|
|
|
|
|
- prections.show()
|
|
|
+ // 将添加了索引的RDD 转化为DataFrame
|
|
|
+ val df2 = sparkSession.createDataFrame(rowRDD, schema)
|
|
|
+ return df2;
|
|
|
+ }
|
|
|
|
|
|
- val results = prections.select("prediction")
|
|
|
|
|
|
|
|
|
|
|
|
- val columns = keyFeature
|
|
|
- val res = results.collect().map(row => {
|
|
|
- row.toSeq.zipWithIndex.map(pair => {
|
|
|
- (columns, pair._1.toString)
|
|
|
- }).toMap.asJava
|
|
|
- }).toList.asJava
|
|
|
- res
|
|
|
- }
|
|
|
- }
|
|
|
|
|
|
}
|