【发布时间】:2020-07-29 12:55:57
【问题描述】:
我创建了一个管道并尝试在 spark 中训练 Kmean 聚类算法,但它失败了,我无法找到确切的错误。这是代码
import org.apache.spark.ml.Pipeline
import org.apache.spark.ml.clustering.KMeans
import org.apache.spark.ml.evaluation.ClusteringEvaluator
import org.apache.spark.ml.feature.{OneHotEncoderEstimator, StringIndexer, VectorAssembler, Normalizer}
import org.apache.spark.{SparkConf, SparkContext}
import org.apache.spark.sql.{SQLContext, SparkSession, functions}
import org.apache.spark.sql.functions._
import org.apache.spark.sql.types.{DoubleType, IntegerType}
val df = spark.read.option("header", "false").option("delimiter", " ").
csv("HMP_Dataset/*").
withColumn("Class" , element_at(reverse(split(input_file_name,"/")),2) ).
withColumn("Source" , element_at(reverse(split(input_file_name,"/")),1)).
withColumnRenamed("_c0","X").withColumnRenamed("_c1","Y").
withColumnRenamed("_c2","Z")
val df2 = df.select(
df.columns.map {
case x @ "X" => df(x).cast(DoubleType).as(x)
case y @ "Y" => df(y).cast(DoubleType).as(y)
case z @ "Z" => df(z).cast(DoubleType).as(z)
case other => df(other)
}: _*
)
val indexer = new StringIndexer().setInputCol("Class").setOutputCol("ClassIndex")
val encoder = new OneHotEncoderEstimator().setInputCols(Array("ClassIndex")) .setOutputCols(Array("CategoryVec"))
val assembler = new VectorAssembler().setInputCols(Array("X","Y","Z")).setOutputCol("Features")
val normalizer = new Normalizer().setInputCol("Features").setOutputCol("feature_Norm")
val pipeline = new Pipeline( ).setStages(Array ( indexer , encoder , assembler , normalizer) )
val model = pipeline.fit(df2).transform(df2)
val train = model.drop("X").drop("Y").drop("Z").drop("Class").drop("Source").drop("ClassIndex").drop("Features")
//model.show()
//train.show()
val kmeans = new KMeans().setFeaturesCol("feature_Norm").setK(2).setSeed(1).setMaxIter(100).fit(train).transform(train)
train 数据框创建成功,但是当我传递给 Kmeans 时,它会抛出错误。错误信息是
Failed to execute user defined function($anonfun$4: (struct<X:double,Y:double,Z:double>) => struct<type:tinyint,size:int,indices:array<int>,values:array<double>>).
我该如何解决这个问题?
【问题讨论】:
-
你能写几行你想读的文件吗?
-
@Chema 这是我要阅读的数据集的链接。
-
我可以看看你的进口数据吗?
-
@Chema 我已经更新了问题。您现在可以导入了。
标签: scala dataframe apache-spark pipeline k-means