| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188 |
- 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.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 fileLocDao: 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 fileLocs = fileLocDao.findById(fileId)
- println(fileLocs)
- val filepath = fileLocs.getLocation
- val filename = fileLocs.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 keyFeature = headerInfoDao.findById(Integer.parseInt(fieldsId.get(fieldsId.size()-1))).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 parquetFile = sparkSession.read.parquet(filepath)
- // parquetFile.select()
- val featuresDF = parquetFile.select(keyFeature, fieldNames: _*).persist()
- featuresDF.schema.foreach(println)
- val strFeatures = featuresDF.schema.filter(_.dataType.equals(StringType)).map(_.name)
- val notNullFeatureDF = featuresDF.na.fill("NA", strFeatures)
- val indexers = strFeatures.map(field => {
- new StringIndexer()
- .setInputCol(field)
- .setOutputCol(field+"_index")
- .fit(notNullFeatureDF)
- })
- //transform String to Indexer
- var tmpDF:DataFrame = null
- if (indexers.nonEmpty) {
- println("Indexers: " + indexers.size)
- val pipeline = new Pipeline().setStages(indexers.toArray[StringIndexerModel])
- println("Pipeline: " + pipeline.getStages.length)
- tmpDF = pipeline.fit(notNullFeatureDF).transform(notNullFeatureDF)
- } else {
- tmpDF = notNullFeatureDF
- }
- tmpDF.show()
- val transformDF = tmpDF.drop(strFeatures: _*)
- val keyFeatureIdx = if(strFeatures.nonEmpty && strFeatures.head.equals(keyFeature)) transformDF.columns.length - strFeatures.size else 0
- val labeledData = transformDF.collect().map(row => {
- val seq = row.toSeq
- val key = seq(keyFeatureIdx)
- val tmp = seq.splitAt(keyFeatureIdx)
- val tail = tmp._1 ++ tmp._2.tail
- val features = tail.zipWithIndex.map(pair =>
- if (pair._1 == null) null else (pair._2 + 1) + ":" + pair._1).filter(_ != null)
- if (key == null) null else key + " " + features.mkString(" ")
- }).filter(_ != null)
- val tmpFileLoc = hdfsServer + config.getId + "_ver2.libsvm"
- //hdfsDao.deleteFileInHdfs(tmpFileLoc, true)
- val labeledDS = sparkSession.createDataset(labeledData)
- labeledDS.repartition(1).write.mode(SaveMode.Overwrite).text(tmpFileLoc)
- log.info("Write to the file: " + tmpFileLoc)
- // Load the data stored in LIBSVM format as a DataFrame.
- val data = sparkSession.read.format("libsvm").load(tmpFileLoc)
- // 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==1) {
- val model1=LogisticRegressionModel.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 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
- }
- }
- }
|