FitTestData.scala 4.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149
  1. package com.example.data
  2. import java.util
  3. import com.example.SparkConnect
  4. import com.example.dao._
  5. import com.example.entity.{Config, Model}
  6. import com.example.util.Json2Object
  7. import org.apache.spark.ml.classification.LogisticRegressionModel
  8. import org.apache.spark.ml.clustering.KMeansModel
  9. import org.apache.spark.ml.feature.{IndexToString, StringIndexer, StringIndexerModel}
  10. import org.apache.spark.ml.{Pipeline, PipelineModel}
  11. import org.apache.spark.rdd.RDD
  12. import org.apache.spark.sql.functions.{col, lit, typedLit}
  13. import org.apache.spark.sql.types.{LongType, StringType, StructField, StructType}
  14. import org.apache.spark.sql.{DataFrame, Row, SaveMode}
  15. import org.springframework.beans.factory.annotation.Autowired
  16. import org.springframework.stereotype.Component
  17. import scala.collection.JavaConverters._
  18. /**
  19. * Created by twb on 2017/5/3.
  20. */
  21. @Component
  22. class FitTestData extends SparkConnect{
  23. @Autowired
  24. val libsvmAdapter: LibsvmAdapter = null
  25. @Autowired
  26. val fileInfoDao: FileInfoDao = null
  27. @Autowired
  28. val hdfsDao: HdfsDao = null
  29. @Autowired
  30. val modelDao: ModelDao = null
  31. @Autowired
  32. val configDao: ConfigDao = null
  33. @Autowired
  34. val headerInfoDao: HeaderInfoDao = null
  35. @Autowired
  36. val dataSetHelper : DataSetHelper = null
  37. @throws(classOf[Exception])
  38. def transformData(fileId: Int, modelId: Int): util.List[util.Map[String, String]] = {
  39. val fileInfo = fileInfoDao.findById(fileId)
  40. log.info("fileId=" + fileId + ", read: "+fileInfo.toString)
  41. val filepath = fileInfo.getLocation
  42. val filename = fileInfo.getFilename
  43. log.info("file loc: " + filepath)
  44. val model: Model = modelDao.findById(modelId)
  45. val config: Config = configDao.findById(model.getConfigId)
  46. val json2Object = new Json2Object
  47. val fieldsId = json2Object.Json2List(config.getFieldIds)//.listIterator()
  48. val iter = fieldsId.listIterator()
  49. //println(fieldsId(0))
  50. //val labelName = fieldsId.remove(fieldsId.size() - 1)
  51. val fieldNames: Array[String] = new Array[String](fieldsId.size()-1)
  52. val labelHeaderInfoId = fieldsId.get(fieldsId.size()-1)
  53. val labelName = headerInfoDao.findById(Integer.parseInt(labelHeaderInfoId)).getFieldName
  54. var i: Int = 0
  55. var flag = false
  56. while(iter.hasNext && !flag){
  57. fieldNames(i) = headerInfoDao.findById(Integer.parseInt(iter.next())).getFieldName
  58. println(fieldNames(i))
  59. i = i+1
  60. if(i>=fieldsId.size()-1)
  61. flag = true
  62. }
  63. //上述代码应该被抽象为一个业务逻辑,给定文件,获取文件的分析数据,需要有一个数据结构支撑!
  64. val spark = sparkSession
  65. import spark.implicits._
  66. val testSet = sparkSession.read.parquet(filepath)
  67. //libsvm文件统一资源定位符。
  68. val libsvmFileUri = dataSetHelper.processTestDataSet(modelId,fileId,fieldNames)
  69. // 默认label 和 特征features
  70. val testDataLibSvm = sparkSession.read.format("libsvm").load(libsvmFileUri)
  71. // 完成预测
  72. val modelPath = model.getModelPath
  73. val trainedModel = PipelineModel.load(modelPath)
  74. val prediction = trainedModel.transform(testDataLibSvm)
  75. prediction.show()
  76. var predictedLabel = prediction.select("prediction")
  77. predictedLabel = dataSetHelper.restoreLabel(modelId, predictedLabel).toDF("prediction_"+labelName)
  78. val res: util.List[util.Map[String, String]] = predictedLabel.collect().map(row => {
  79. row.toSeq.zipWithIndex.map(pair => {
  80. (labelName, pair._1.toString)
  81. }).toMap.asJava
  82. }).toList.asJava
  83. testSet.show()
  84. val testSetWithId = addId(testSet)
  85. testSetWithId.show()
  86. predictedLabel.show()
  87. val predictedLabelWithId = addId(predictedLabel)
  88. predictedLabelWithId.show()
  89. var predictTable:DataFrame = testSetWithId.join(predictedLabelWithId,"id");
  90. predictTable.show()
  91. predictTable = predictTable.orderBy("id")
  92. predictTable.show()
  93. res
  94. }
  95. private def addId(dataFrame:DataFrame):DataFrame={
  96. val schema: StructType = dataFrame.schema.add(StructField("id", LongType))
  97. // DataFrame转RDD 然后调用 zipWithIndex
  98. //zipWithIndex返回的是元组数组:Array((Sunday,0), (Monday,1), ...
  99. //每一元组,第一个为原来的Row,第二项是index;
  100. val dfRDD: RDD[(Row, Long)] = dataFrame.rdd.zipWithIndex()
  101. //合并数组中的元组,第二项需要转换为Row
  102. val rowRDD: RDD[Row] = dfRDD.map(tp => Row.merge(tp._1, Row(tp._2)))
  103. // 将添加了索引的RDD 转化为DataFrame
  104. val df2 = sparkSession.createDataFrame(rowRDD, schema)
  105. return df2;
  106. }
  107. }