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

如何在Python/PySpark中高效计算带权重的列中位数(不扩展行)

加权中位数计算(不扩展行)

Python(Pandas)实现

核心思路是通过排序+累计权重定位找到中位数,无需扩展行,适合中小数据量场景:

import pandas as pd

# 构造示例DataFrame
df = pd.DataFrame({
    'col1': [100, 200, 300, 400],
    'col2': [3, 2, 4, 1]
})

# 1. 按col1排序,确保数据有序(中位数依赖有序序列)
df_sorted = df.sort_values('col1')
# 2. 计算累计权重
df_sorted['cum_weight'] = df_sorted['col2'].cumsum()
# 3. 计算总权重和中位数阈值
total_weight = df_sorted['col2'].sum()
median_threshold = total_weight / 2

# 4. 根据总权重奇偶性计算加权中位数
if total_weight % 2 == 1:
    # 奇数权重:取第一个累计权重超过阈值的col1值
    weighted_median = df_sorted[df_sorted['cum_weight'] > median_threshold]['col1'].iloc[0]
else:
    # 偶数权重:取两个临界位置的col1值的平均值
    lower_val = df_sorted[df_sorted['cum_weight'] >= median_threshold - 0.5]['col1'].iloc[0]
    upper_val = df_sorted[df_sorted['cum_weight'] >= median_threshold + 0.5]['col1'].iloc[0]
    weighted_median = (lower_val + upper_val) / 2

print(weighted_median)  # 输出:250

PySpark实现

针对大数据场景,用窗口函数计算累计权重,完全避免行扩展,性能更优:

from pyspark.sql import SparkSession
from pyspark.sql import functions as F
from pyspark.sql.window import Window

# 初始化Spark会话
spark = SparkSession.builder.appName("WeightedMedian").getOrCreate()

# 构造示例DataFrame
data = [(100, 3), (200, 2), (300, 4), (400, 1)]
df = spark.createDataFrame(data, ["col1", "col2"])

# 1. 计算总权重和中位数阈值
total_weight = df.agg(F.sum("col2")).first()[0]
median_threshold = total_weight / 2

# 2. 定义窗口:按col1排序,计算从第一行到当前行的累计权重
window_spec = Window.orderBy("col1").rowsBetween(Window.unboundedPreceding, Window.currentRow)
df_with_cum = df.withColumn("cum_weight", F.sum("col2").over(window_spec))

# 3. 根据总权重奇偶性计算加权中位数
if total_weight % 2 == 1:
    # 奇数权重:筛选第一个累计权重超过阈值的col1
    weighted_median = df_with_cum.filter(F.col("cum_weight") > median_threshold).select("col1").first()[0]
else:
    # 偶数权重:获取两个临界值并计算平均
    lower_val = df_with_cum.filter(F.col("cum_weight") >= median_threshold - 0.5).select("col1").limit(1).first()[0]
    upper_val = df_with_cum.filter(F.col("cum_weight") >= median_threshold + 0.5).select("col1").limit(1).first()[0]
    weighted_median = (lower_val + upper_val) / 2

print(weighted_median)  # 输出:250

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 17:17:38