FitTestData.scala 4.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150
  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.{StringIndexer, StringIndexerModel}
  10. import org.apache.spark.ml.{Pipeline, PipelineModel}
  11. import org.apache.spark.sql.functions.lit
  12. import org.apache.spark.sql.types.StringType
  13. import org.apache.spark.sql.{DataFrame, SaveMode}
  14. import org.springframework.beans.factory.annotation.Autowired
  15. import org.springframework.stereotype.Component
  16. import scala.collection.JavaConverters._
  17. /**
  18. * Created by twb on 2017/5/3.
  19. */
  20. @Component
  21. class FitTestData extends SparkConnect{
  22. @Autowired
  23. val libsvmAdapter: LibsvmAdapter = null
  24. @Autowired
  25. val fileInfoDao: FileInfoDao = null
  26. @Autowired
  27. val hdfsDao: HdfsDao = null
  28. @Autowired
  29. val modelDao: ModelDao = null
  30. @Autowired
  31. val configDao: ConfigDao = null
  32. @Autowired
  33. val headerInfoDao: HeaderInfoDao = null
  34. @throws(classOf[Exception])
  35. def transformData(fileId: Int, modelId: Int): util.List[util.Map[String, String]] = {
  36. val fileInfo = fileInfoDao.findById(fileId)
  37. log.info("fileId=" + fileId + ", read: "+fileInfo.toString)
  38. val filepath = fileInfo.getLocation
  39. val filename = fileInfo.getFilename
  40. log.info("file loc: " + filepath)
  41. val model: Model = modelDao.findById(modelId)
  42. val config: Config = configDao.findById(model.getConfigId)
  43. val json2Object = new Json2Object
  44. val fieldsId = json2Object.Json2List(config.getFieldIds)//.listIterator()
  45. val iter = fieldsId.listIterator()
  46. //println(fieldsId(0))
  47. //val keyFeature = fieldsId.remove(fieldsId.size() - 1)
  48. val fieldNames: Array[String] = new Array[String](fieldsId.size()-1)
  49. val labelHeaderInfoId = fieldsId.get(fieldsId.size()-1)
  50. val keyFeature = headerInfoDao.findById(Integer.parseInt(labelHeaderInfoId)).getFieldName
  51. var i: Int = 0
  52. var flag = false
  53. while(iter.hasNext && !flag){
  54. fieldNames(i) = headerInfoDao.findById(Integer.parseInt(iter.next())).getFieldName
  55. println(fieldNames(i))
  56. i = i+1
  57. if(i>=fieldsId.size()-1)
  58. flag = true
  59. }
  60. val spark = sparkSession
  61. import spark.implicits._
  62. val dataFrameFromParquet = sparkSession.read.parquet(filepath)
  63. //此处需要做一个适配,若文件此时没有label列,就先加上
  64. var dataFrameForPredict:DataFrame = dataFrameFromParquet
  65. if(!dataFrameFromParquet.columns.toList.contains(keyFeature)){
  66. val labelHeaderInfo = headerInfoDao.findById(Integer.parseInt(labelHeaderInfoId))
  67. val valueInfo = labelHeaderInfo.getValueInfo
  68. val fieldType = labelHeaderInfo.getFieldType
  69. val valueArray = json2Object.Json2List(valueInfo)
  70. val labelDefaultValueStr = valueArray.get(0)
  71. val labelDefaultValue = fieldType match {
  72. case "string" => labelDefaultValueStr
  73. case "int" => Integer.parseInt(labelDefaultValueStr)
  74. case "double" => labelDefaultValueStr.toDouble
  75. case _ => log.error("类型出错")
  76. }
  77. dataFrameForPredict = dataFrameFromParquet.withColumn(keyFeature,lit(labelDefaultValue))
  78. }
  79. //添加部分到此结束
  80. //libsvm文件统一资源定位符。
  81. val libsvmFileUri = libsvmAdapter.dataFrame2Libsvm(dataFrameForPredict,fileId,keyFeature,fieldNames)
  82. // Load the data stored in LIBSVM format as a DataFrame.
  83. val data = sparkSession.read.format("libsvm").load(libsvmFileUri)
  84. // Index labels, adding metadata to the label column.
  85. // Fit on whole dataset to include all labels in index.
  86. //modelPath : hdfsServer + "_" + modelId + ".model"
  87. val modelPath = model.getModelPath
  88. if (model.getModelTypeId==5){
  89. val model1 = KMeansModel.load(modelPath)
  90. val prections = model1.transform(data)
  91. prections.show()
  92. val results = prections.select("predictedLabel")
  93. val columns = keyFeature
  94. val res = results.collect().map(row => {
  95. row.toSeq.zipWithIndex.map(pair => {
  96. (columns, pair._1.toString)
  97. }).toMap.asJava
  98. }).toList.asJava
  99. res
  100. }else{
  101. val model1 = PipelineModel.load(modelPath)
  102. val prections = model1.transform(data)
  103. prections.show()
  104. val results = prections.select("predictedLabel")
  105. val columns = keyFeature
  106. val res = results.collect().map(row => {
  107. row.toSeq.zipWithIndex.map(pair => {
  108. (columns, pair._1.toString)
  109. }).toMap.asJava
  110. }).toList.asJava
  111. res
  112. }
  113. }
  114. }