| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150 |
- package com.example.data
- import java.util
- 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.{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.springframework.beans.factory.annotation.Autowired
- import org.springframework.stereotype.Component
- import scala.collection.JavaConverters._
- /**
- * Created by twb on 2017/5/3.
- */
- @Component
- class FitTestData extends SparkConnect{
- @Autowired
- val libsvmAdapter: LibsvmAdapter = null
- @Autowired
- val fileInfoDao: FileInfoDao = null
- @Autowired
- val hdfsDao: HdfsDao = null
- @Autowired
- val modelDao: ModelDao = null
- @Autowired
- val configDao: ConfigDao = null
- @Autowired
- val headerInfoDao: HeaderInfoDao = null
- @throws(classOf[Exception])
- def transformData(fileId: Int, modelId: Int): util.List[util.Map[String, String]] = {
- val fileInfo = fileInfoDao.findById(fileId)
- log.info("fileId=" + fileId + ", read: "+fileInfo.toString)
- val filepath = fileInfo.getLocation
- val filename = fileInfo.getFilename
- log.info("file loc: " + filepath)
- val model: Model = modelDao.findById(modelId)
- val config: Config = configDao.findById(model.getConfigId)
- val json2Object = new Json2Object
- val fieldsId = json2Object.Json2List(config.getFieldIds)//.listIterator()
- val iter = fieldsId.listIterator()
- //println(fieldsId(0))
- //val keyFeature = 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
- var i: Int = 0
- var flag = false
- while(iter.hasNext && !flag){
- fieldNames(i) = headerInfoDao.findById(Integer.parseInt(iter.next())).getFieldName
- println(fieldNames(i))
- i = i+1
- if(i>=fieldsId.size()-1)
- flag = true
- }
- val spark = sparkSession
- import spark.implicits._
- val dataFrameFromParquet = sparkSession.read.parquet(filepath)
- //此处需要做一个适配,若文件此时没有label列,就先加上
- var dataFrameForPredict:DataFrame = dataFrameFromParquet
- if(!dataFrameFromParquet.columns.toList.contains(keyFeature)){
- val labelHeaderInfo = headerInfoDao.findById(Integer.parseInt(labelHeaderInfoId))
- val valueInfo = labelHeaderInfo.getValueInfo
- val fieldType = labelHeaderInfo.getFieldType
- val valueArray = json2Object.Json2List(valueInfo)
- val labelDefaultValueStr = valueArray.get(0)
- val labelDefaultValue = fieldType match {
- case "string" => labelDefaultValueStr
- case "int" => Integer.parseInt(labelDefaultValueStr)
- case "double" => labelDefaultValueStr.toDouble
- case _ => log.error("类型出错")
- }
- dataFrameForPredict = dataFrameFromParquet.withColumn(keyFeature,lit(labelDefaultValue))
- }
- //添加部分到此结束
- //libsvm文件统一资源定位符。
- val libsvmFileUri = libsvmAdapter.dataFrame2Libsvm(dataFrameForPredict,fileId,keyFeature,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.
- //modelPath : hdfsServer + "_" + modelId + ".model"
- val modelPath = model.getModelPath
- if (model.getModelTypeId==5){
- val model1 = KMeansModel.load(modelPath)
- val prections = model1.transform(data)
- prections.show()
- val results = prections.select("predictedLabel")
- val columns = keyFeature
- val res = results.collect().map(row => {
- row.toSeq.zipWithIndex.map(pair => {
- (columns, pair._1.toString)
- }).toMap.asJava
- }).toList.asJava
- res
- }else{
- val model1 = PipelineModel.load(modelPath)
- val prections = model1.transform(data)
- prections.show()
- val results = prections.select("predictedLabel")
- val columns = keyFeature
- val res = results.collect().map(row => {
- row.toSeq.zipWithIndex.map(pair => {
- (columns, pair._1.toString)
- }).toMap.asJava
- }).toList.asJava
- res
- }
- }
- }
|