FitTestData.scala 5.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188
  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.types.StringType
  12. import org.apache.spark.sql.{DataFrame, SaveMode}
  13. import org.springframework.beans.factory.annotation.Autowired
  14. import org.springframework.stereotype.Component
  15. import scala.collection.JavaConverters._
  16. /**
  17. * Created by twb on 2017/5/3.
  18. */
  19. @Component
  20. class FitTestData extends SparkConnect{
  21. @Autowired
  22. val fileLocDao: FileInfoDao = null
  23. @Autowired
  24. val hdfsDao: HdfsDao = null
  25. @Autowired
  26. val modelDao: ModelDao = null
  27. @Autowired
  28. val configDao: ConfigDao = null
  29. @Autowired
  30. val headerInfoDao: HeaderInfoDao = null
  31. @throws(classOf[Exception])
  32. def transformData(fileId: Int, modelId: Int): util.List[util.Map[String, String]] = {
  33. val fileLocs = fileLocDao.findById(fileId)
  34. println(fileLocs)
  35. val filepath = fileLocs.getLocation
  36. val filename = fileLocs.getFilename
  37. log.info("file loc: " + filepath)
  38. val model: Model = modelDao.findById(modelId)
  39. val config: Config = configDao.findById(model.getConfigId)
  40. val json2Object = new Json2Object
  41. val fieldsId = json2Object.Json2List(config.getFieldIds)//.listIterator()
  42. val iter = fieldsId.listIterator()
  43. //println(fieldsId(0))
  44. //val keyFeature = fieldsId.remove(fieldsId.size() - 1)
  45. val fieldNames: Array[String] = new Array[String](fieldsId.size()-1)
  46. val keyFeature = headerInfoDao.findById(Integer.parseInt(fieldsId.get(fieldsId.size()-1))).getFieldName
  47. var i: Int = 0
  48. var flag = false
  49. while(iter.hasNext && !flag){
  50. fieldNames(i) = headerInfoDao.findById(Integer.parseInt(iter.next())).getFieldName
  51. println(fieldNames(i))
  52. i = i+1
  53. if(i>=fieldsId.size()-1)
  54. flag = true
  55. }
  56. val spark = sparkSession
  57. import spark.implicits._
  58. val parquetFile = sparkSession.read.parquet(filepath)
  59. // parquetFile.select()
  60. val featuresDF = parquetFile.select(keyFeature, fieldNames: _*).persist()
  61. featuresDF.schema.foreach(println)
  62. val strFeatures = featuresDF.schema.filter(_.dataType.equals(StringType)).map(_.name)
  63. val notNullFeatureDF = featuresDF.na.fill("NA", strFeatures)
  64. val indexers = strFeatures.map(field => {
  65. new StringIndexer()
  66. .setInputCol(field)
  67. .setOutputCol(field+"_index")
  68. .fit(notNullFeatureDF)
  69. })
  70. //transform String to Indexer
  71. var tmpDF:DataFrame = null
  72. if (indexers.nonEmpty) {
  73. println("Indexers: " + indexers.size)
  74. val pipeline = new Pipeline().setStages(indexers.toArray[StringIndexerModel])
  75. println("Pipeline: " + pipeline.getStages.length)
  76. tmpDF = pipeline.fit(notNullFeatureDF).transform(notNullFeatureDF)
  77. } else {
  78. tmpDF = notNullFeatureDF
  79. }
  80. tmpDF.show()
  81. val transformDF = tmpDF.drop(strFeatures: _*)
  82. val keyFeatureIdx = if(strFeatures.nonEmpty && strFeatures.head.equals(keyFeature)) transformDF.columns.length - strFeatures.size else 0
  83. val labeledData = transformDF.collect().map(row => {
  84. val seq = row.toSeq
  85. val key = seq(keyFeatureIdx)
  86. val tmp = seq.splitAt(keyFeatureIdx)
  87. val tail = tmp._1 ++ tmp._2.tail
  88. val features = tail.zipWithIndex.map(pair =>
  89. if (pair._1 == null) null else (pair._2 + 1) + ":" + pair._1).filter(_ != null)
  90. if (key == null) null else key + " " + features.mkString(" ")
  91. }).filter(_ != null)
  92. val tmpFileLoc = hdfsServer + config.getId + "_ver2.libsvm"
  93. //hdfsDao.deleteFileInHdfs(tmpFileLoc, true)
  94. val labeledDS = sparkSession.createDataset(labeledData)
  95. labeledDS.repartition(1).write.mode(SaveMode.Overwrite).text(tmpFileLoc)
  96. log.info("Write to the file: " + tmpFileLoc)
  97. // Load the data stored in LIBSVM format as a DataFrame.
  98. val data = sparkSession.read.format("libsvm").load(tmpFileLoc)
  99. // Index labels, adding metadata to the label column.
  100. // Fit on whole dataset to include all labels in index.
  101. //modelPath : hdfsServer + "_" + modelId + ".model"
  102. val modelPath = model.getModelPath
  103. if (model.getModelTypeId==1) {
  104. val model1=LogisticRegressionModel.load(modelPath)
  105. val prections = model1.transform(data)
  106. prections.show()
  107. val results = prections.select("predictedLabel")
  108. val columns = keyFeature
  109. val res = results.collect().map(row => {
  110. row.toSeq.zipWithIndex.map(pair => {
  111. (columns, pair._1.toString)
  112. }).toMap.asJava
  113. }).toList.asJava
  114. res
  115. }else if (model.getModelTypeId==5){
  116. val model1 = KMeansModel.load(modelPath)
  117. val prections = model1.transform(data)
  118. prections.show()
  119. val results = prections.select("predictedLabel")
  120. val columns = keyFeature
  121. val res = results.collect().map(row => {
  122. row.toSeq.zipWithIndex.map(pair => {
  123. (columns, pair._1.toString)
  124. }).toMap.asJava
  125. }).toList.asJava
  126. res
  127. }else{
  128. val model1 = PipelineModel.load(modelPath)
  129. val prections = model1.transform(data)
  130. prections.show()
  131. val results = prections.select("predictedLabel")
  132. val columns = keyFeature
  133. val res = results.collect().map(row => {
  134. row.toSeq.zipWithIndex.map(pair => {
  135. (columns, pair._1.toString)
  136. }).toMap.asJava
  137. }).toList.asJava
  138. res
  139. }
  140. }
  141. }