|
|
@@ -7,7 +7,7 @@ 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.feature.{IndexToString, StringIndexer, StringIndexerModel}
|
|
|
import org.apache.spark.ml.{Pipeline, PipelineModel}
|
|
|
import org.apache.spark.sql.functions.lit
|
|
|
import org.apache.spark.sql.types.StringType
|
|
|
@@ -39,6 +39,8 @@ class FitTestData extends SparkConnect{
|
|
|
@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]] = {
|
|
|
@@ -80,32 +82,8 @@ class FitTestData extends SparkConnect{
|
|
|
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)
|
|
|
+ val libsvmFileUri = dataSetHelper.processTestDataSet(modelId,fileId,fieldNames)
|
|
|
|
|
|
// Load the data stored in LIBSVM format as a DataFrame.
|
|
|
val data = sparkSession.read.format("libsvm").load(libsvmFileUri)
|
|
|
@@ -121,7 +99,7 @@ class FitTestData extends SparkConnect{
|
|
|
|
|
|
prections.show()
|
|
|
|
|
|
- val results = prections.select("predictedLabel")
|
|
|
+ val results = prections.select("prediction")
|
|
|
val columns = keyFeature
|
|
|
val res = results.collect().map(row => {
|
|
|
row.toSeq.zipWithIndex.map(pair => {
|
|
|
@@ -136,7 +114,10 @@ class FitTestData extends SparkConnect{
|
|
|
|
|
|
prections.show()
|
|
|
|
|
|
- val results = prections.select("predictedLabel")
|
|
|
+ val results = prections.select("prediction")
|
|
|
+
|
|
|
+
|
|
|
+
|
|
|
val columns = keyFeature
|
|
|
val res = results.collect().map(row => {
|
|
|
row.toSeq.zipWithIndex.map(pair => {
|