You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在Spark中无需collectAsList获取数据集列的中位数?

Spark 分布式计算中位数方案(避免拉取全量数据)

核心思路是利用窗口函数给排序后的每行分配行号,结合总行数定位中位数所在行,全程在分布式环境执行,无需将全量数据拉到Driver端。以下是多语言实现:

Java 实现

import org.apache.spark.sql.Dataset;
import org.apache.spark.sql.Row;
import org.apache.spark.sql.expressions.Window;
import org.apache.spark.sql.expressions.WindowSpec;
import static org.apache.spark.sql.functions.*;

public String getMedian(Dataset<Row> dataset, String column) {
    // 过滤空值并获取有效数据行数
    long totalCount = dataset.where(col(column).isNotNull()).count();
    if (totalCount == 0) {
        return null; // 处理空数据集场景
    }

    // 计算中位数对应的1-based行号
    long medianRowNum;
    if (totalCount % 2 == 1) {
        medianRowNum = (totalCount + 1) / 2; // 奇数个元素取中间行
    } else {
        // 偶数个元素取中间右侧行,如需左侧可改为 totalCount / 2
        medianRowNum = totalCount / 2 + 1;
    }

    // 定义排序窗口
    WindowSpec windowSpec = Window.orderBy(col(column));

    // 分配行号、过滤中位数行、提取结果
    return dataset.where(col(column).isNotNull())
            .withColumn("row_num", row_number().over(windowSpec))
            .where(col("row_num").equalTo(medianRowNum))
            .select(column)
            .head() // 仅提取结果行,不拉取全量数据
            .getString(0);
}

Scala 实现

import org.apache.spark.sql.{Dataset, Row}
import org.apache.spark.sql.expressions.Window
import org.apache.spark.sql.functions._

def getMedian(dataset: Dataset[Row], column: String): String = {
    val filteredDs = dataset.where(col(column).isNotNull)
    val totalCount = filteredDs.count()
    if (totalCount == 0) return null

    val medianRowNum = if (totalCount % 2 == 1) (totalCount + 1) / 2 else totalCount / 2 + 1

    val windowSpec = Window.orderBy(col(column))

    filteredDs.withColumn("row_num", row_number().over(windowSpec))
        .where(col("row_num") === medianRowNum)
        .select(column)
        .head()
        .getString(0)
}

PySpark 实现

from pyspark.sql import Window
from pyspark.sql.functions import row_number, col

def get_median(dataset, column):
    filtered_ds = dataset.where(col(column).isNotNull())
    total_count = filtered_ds.count()
    if total_count == 0:
        return None
    
    if total_count % 2 == 1:
        median_row_num = (total_count + 1) // 2
    else:
        median_row_num = total_count // 2 + 1  # 如需中间左侧行,改为 total_count//2
    
    window_spec = Window.orderBy(col(column))
    
    return filtered_ds.withColumn("row_num", row_number().over(window_spec))\
        .where(col("row_num") == median_row_num)\
        .select(column)\
        .head()[0]

关键说明

  • 规避全量拉取:用head()仅提取结果集中的唯一行,替代collectAsList(),避免Driver端内存溢出风险。
  • 修正原逻辑问题:原代码对奇数元素的中位数计算有误(示例中会错误取到Patrick),上述实现中奇数场景取中间行,符合示例预期的Mel。
  • 灵活适配偶数场景:可根据需求调整偶数元素的中位数取值(中间左侧/右侧)。

内容的提问来源于stack exchange,提问作者SSSOF

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.25 22:03:31