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; } }