最简单的方法(需要 Spark 2.0.1+ 而不是精确的中位数)
正如 cmets 中提到的第一个问题 Find median in Spark SQL for double datatype columns 中所述,我们可以使用 percentile_approx 来计算 Spark 2.0.1+ 的中值。要将其应用于 Apache Spark 中的分组数据,查询应如下所示:
val df = Seq(("A", 0.0), ("A", 0.0), ("A", 1.0), ("A", 1.0), ("A", 1.0), ("A", 1.0), ("B", 0.0), ("B", 1.0), ("B", 1.0)).toDF("id", "num")
df.createOrReplaceTempView("df")
spark.sql("select id, percentile_approx(num, 0.5) as median from df group by id order by id").show()
输出为:
+---+------+
| id|median|
+---+------+
| A| 1.0|
| B| 1.0|
+---+------+
也就是说,这是一个近似值(而不是每个问题的精确中位数)。
计算分组数据的准确中位数
有多种方法,所以我相信 SO 中的其他人可以提供更好或更有效的示例。但这里有一段代码 sn-p 计算 Spark 中分组数据的中位数(在 Spark 1.6 和 Spark 2.1 中验证):
import org.apache.spark.SparkContext._
val rdd: RDD[(String, Double)] = sc.parallelize(Seq(("A", 1.0), ("A", 0.0), ("A", 1.0), ("A", 1.0), ("A", 0.0), ("A", 1.0), ("B", 0.0), ("B", 1.0), ("B", 1.0)))
// Scala median function
def median(inputList: List[Double]): Double = {
val count = inputList.size
if (count % 2 == 0) {
val l = count / 2 - 1
val r = l + 1
(inputList(l) + inputList(r)).toDouble / 2
} else
inputList(count / 2).toDouble
}
// Sort the values
val setRDD = rdd.groupByKey()
val sortedListRDD = setRDD.mapValues(_.toList.sorted)
// Output DataFrame of id and median
sortedListRDD.map(m => {
(m._1, median(m._2))
}).toDF("id", "median_of_num").show()
输出为:
+---+-------------+
| id|median_of_num|
+---+-------------+
| A| 1.0|
| B| 1.0|
+---+-------------+
我应该指出一些警告,因为这可能不是最有效的实现:
- 目前使用的
groupByKey 性能不是很好。您可能希望将其更改为 reduceByKey(更多信息请访问 Avoid GroupByKey)
- 使用 Scala 函数计算
median。
这种方法应该适用于少量数据,但如果每个键都有数百万行,建议使用 Spark 2.0.1+ 并使用 percentile_approx 方法。