| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149 |
- 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.{IndexToString, StringIndexer, StringIndexerModel}
- import org.apache.spark.ml.{Pipeline, PipelineModel}
- 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
- 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
- @Autowired
- val dataSetHelper : DataSetHelper = 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 labelName = fieldsId.remove(fieldsId.size() - 1)
- val fieldNames: Array[String] = new Array[String](fieldsId.size()-1)
- val labelHeaderInfoId = fieldsId.get(fieldsId.size()-1)
- val labelName = 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 testSet = sparkSession.read.parquet(filepath)
- //libsvm文件统一资源定位符。
- val libsvmFileUri = dataSetHelper.processTestDataSet(modelId,fileId,fieldNames)
- // 默认label 和 特征features
- val testDataLibSvm = sparkSession.read.format("libsvm").load(libsvmFileUri)
- // 完成预测
- val modelPath = model.getModelPath
- 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
- }
- 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()
- //合并数组中的元组,第二项需要转换为Row
- val rowRDD: RDD[Row] = dfRDD.map(tp => Row.merge(tp._1, Row(tp._2)))
- // 将添加了索引的RDD 转化为DataFrame
- val df2 = sparkSession.createDataFrame(rowRDD, schema)
- return df2;
- }
- }
|