Spark/Pandas中如何同步计算数组列每行的中位数?
计算DataFrame数组列的每行中位数(Spark/Pandas高效实现)
Spark 实现(无UDF批量处理)
Spark内置函数可直接实现批量计算,无需逐行迭代的UDF,效率更优:
from pyspark.sql import SparkSession from pyspark.sql.functions import sort_array, size, element_at, when, col spark = SparkSession.builder.appName("array_median").getOrCreate() # 构造示例DataFrame data = [([1,3,4],), ([4,3,5],)] df = spark.createDataFrame(data, ["arr_col"]) # 计算每行数组的中位数 df_result = df.withColumn("sorted_arr", sort_array(col("arr_col"))) \ .withColumn("arr_len", size(col("sorted_arr"))) \ .withColumn("median", when( col("arr_len") % 2 == 1, element_at(col("sorted_arr"), (col("arr_len") + 1) // 2) ).otherwise( (element_at(col("sorted_arr"), col("arr_len") // 2) + element_at(col("sorted_arr"), col("arr_len") // 2 + 1)) / 2 )) \ .select("median").withColumnRenamed("median", "Result") df_result.show()
核心逻辑:
- 用
sort_array对每行数组排序 size获取数组长度,判断奇偶性- 奇数长度取中间位置元素,偶数长度取中间两个元素的平均值
完全基于Spark内置函数,支持分布式批量处理,避免UDF的逐行开销。
Pandas 实现(向量化批量计算)
利用numpy的向量化操作替代逐行apply,大幅提升效率:
import pandas as pd import numpy as np # 构造示例DataFrame data = {"arr_col": [[1,3,4], [4,3,5]]} df = pd.DataFrame(data) # 批量计算每行中位数 df["Result"] = np.median(np.array(df["arr_col"].tolist()), axis=1) # 输出结果 print(df[["Result"]])
核心逻辑:
- 将数组列转为numpy二维数组
- 调用
np.median并指定axis=1,一次性计算所有行的中位数
numpy的向量化操作比逐行apply(np.median)效率高得多,尤其适合大数据量场景。
内容的提问来源于stack exchange,提问作者Barushkish
相关产品推荐
相关产品推荐

