【发布时间】:2021-08-16 15:04:09
【问题描述】:
我知道 KMeansModel transform 将输出作为数据集给我们,并且输出数据帧的预测列说明了哪一列是 _c0、_c1、features 和 prediction。
但是,我也想知道这个数据框中每个特征的每个聚类中心。
我如何使用 java 来做到这一点?我希望我的代码是纯 java 的(不在 scala 中,甚至基于 scala 的 spark)。
+-----------+----------+
| features|prediction|
+-----------+----------+
| [4.0,53.0]| 2|
| [5.0,63.0]| 3|
|[10.0,59.0]| 2|
|[13.0,49.0]| 0|
|[12.0,88.0]| 1|
|[12.0,88.0]| 1|
|[18.0,61.0]| 2|
+-----------+----------+
预期结果:
+-----------+----------+----------------+
| features|prediction|clusterCenter |
+-----------+----------+-------------+--+
| [4.0,53.0]| 1| [20.15,64.95] |
| [5.0,63.0]| 0| [43.91,146.04 |
|[10.0,59.0]| 2| [20.4,68] |
|[13.0,49.0]| 0| [43.91,146.04] |
|[12.0,88.0]| 1| [20.15,64.95] |
|[12.0,88.0]| 2| [20.4,68] |
|[18.0,61.0]| 3| [98.176,114.88]|
+-----------+----------+----------------+
这里是一些测试代码sn-p
List<Row> dataset = Arrays.asList(
RowFactory.create(Vectors.dense(4,53)),
RowFactory.create(Vectors.dense(5,63)),
RowFactory.create(Vectors.dense(10,59)),
RowFactory.create(Vectors.dense(13,49)),
RowFactory.create(Vectors.dense( 12,88)),
RowFactory.create(Vectors.dense(12,88)),
RowFactory.create(Vectors.dense(18,61))
);
StructType schema = new StructType(new StructField[]{
new StructField("features", new VectorUDT(), false, Metadata.empty()),
});
Dataset<Row> df = sc.createDataFrame(dataset, schema);
KMeans kMeans = new KMeans().setK(4).setMaxIter(10);
KMeansModel model = kMeans.fit(df);
Dataset<Row> predict = model.transform(df);
predict.show();
StructType centroidSchema = new StructType(new StructField[]{
new StructField("x", DataTypes.StringType, false, Metadata.empty()),
new StructField("y", DataTypes.StringType, false, Metadata.empty())
});
Dataset<Row> centroid = sc.createDataFrame(jsc.parallelize(model.clusterCenters()).map(s -> {
String[] row = s.toString().replace("[","").replace("]","").split(",");
return RowFactory.create((Object[]) row);
}), centroidSchema);
centroid = centroid .withColumn("x", centroid .col("x").cast("Double"));
centroid = centroid .withColumn("y", centroid .col("y").cast("Double"));
centroid.show();
【问题讨论】:
标签: java apache-spark machine-learning apache-spark-sql