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

PySpark中如何并行计算DataFrame的多种聚合统计量?

PySpark大型DataFrame自定义汇总统计量计算方案对比

测试数据集构造

from pyspark.sql import SparkSession
from pyspark.sql.dataframe import DataFrame
from pyspark.sql.types import DataType, NumericType, DateType, TimestampType
import pyspark.sql.types as t
import pyspark.sql.functions as f
from datetime import datetime

spark = (
    SparkSession.builder
    .appName("pyspark")
    .master("local[*]")
    .getOrCreate()
)
dd = [
    ("Alice", 18.0, datetime(2022, 1, 1)),
    ("Bob", None, datetime(2022, 2, 1)),
    ("Mark", 33.0, None),
    (None, 80.0, datetime(2022, 4, 1)),
]
schema = t.StructType(
    [
        t.StructField("T", t.StringType()),
        t.StructField("C", t.DoubleType()),
        t.StructField("D", t.DateType()),
    ]
)
df = spark.createDataFrame(dd, schema)

需求:对所有列计算missing counts、stddev、max和min,要求并行执行,最终将结果整理为JSON格式。

方案一:单Select查询

通过编写一个大的Select查询,交由Spark引擎完成并行计算:

from typing import List, Tuple

def df_dtypes(df: DataFrame) -> List[Tuple[str, DataType]]:
    """
    Like df.dtypes attribute of Spark DataFrame, but returning DataType objects instead
    of strings.
    """
    return [(str(f.name), f.dataType) for f in df.schema.fields]

def get_missing(df: DataFrame) -> Tuple:
    suffix = "__missing"
    result = (
        *(
            (
                f.count(
                    f.when(
                        (f.isnan(c) | f.isnull(c)),
                        c,
                    )
                )
                / f.count("*")
                * 100
                if isinstance(t, NumericType)  # isnan only works for numeric types
                else f.count(
                    f.when(
                        f.isnull(c),
                        c,
                    )
                )
                / f.count("*")
                * 100
            )
            .cast("double")
            .alias(c + suffix)
            for c, t in df_dtypes(df)
        ),
    )

    return result

def get_min(df: DataFrame) -> Tuple:
    suffix = "__min"
    result = (
        *(
            (f.min(c) if isinstance(t, (NumericType, DateType, TimestampType)) else f.lit(None))
            .cast(t)
            .alias(c + suffix)
            for c, t in df_dtypes(df)
        ),
    )
    return result

def get_max(df: DataFrame) -> Tuple:
    suffix = "__max"
    result = (
        *(
            (f.max(c) if isinstance(t, (NumericType, DateType, TimestampType)) else f.lit(None))
            .cast(t)
            .alias(c + suffix)
            for c, t in df_dtypes(df)
        ),
    )
    return result

def get_std(df: DataFrame) -> Tuple:
    suffix = "__std"
    result = (
        *(
            (f.stddev(c) if isinstance(t, NumericType) else f.lit(None)).cast(t).alias(c + suffix)
            for c, t in df_dtypes(df)
        ),
    )
    return result

# build the big query
query = get_min(df) + get_max(df) + get_missing(df) + get_std(df)

# run the job
df.select(*query).show()

疑问:该方案借助Spark内部机制实现并行,但生成带后缀的大量列是否会成为性能瓶颈?

方案二:使用线程

通过Python线程实现各类计算的并发执行:

from pyspark import InheritableThread
from queue import Queue
from typing import List, Tuple

def df_dtypes(df: DataFrame) -> List[Tuple[str, DataType]]:
    """
    Like df.dtypes attribute of Spark DataFrame, but returning DataType objects instead
    of strings.
    """
    return [(str(f.name), f.dataType) for f in df.schema.fields]

def get_min(df: DataFrame, q: Queue) -> None:
    result = df.select(
        f.lit("min").alias("summary"),
        *(
            (f.min(c) if isinstance(t, (NumericType, DateType, TimestampType)) else f.lit(None))
            .cast(t)
            .alias(c)
            for c, t in df_dtypes(df)
        ),
    ).collect()
    q.put(result)


def get_max(df: DataFrame, q: Queue) -> None:
    result = df.select(
        f.lit("max").alias("summary"),
        *(
            (f.max(c) if isinstance(t, (NumericType, DateType, TimestampType)) else f.lit(None))
            .cast(t)
            .alias(c)
            for c, t in df_dtypes(df)
        ),
    ).collect()
    q.put(result)


def get_std(df: DataFrame, q: Queue) -> None:
    result = df.select(
        f.lit("std").alias("summary"),
        *(
            (f.stddev(c) if isinstance(t, NumericType) else f.lit(None)).cast(t).alias(c)
            for c, t in df_dtypes(df)
        ),
    ).collect()
    q.put(result)

def get_missing(df: DataFrame, q: Queue) -> None:
    result = df.select(
        f.lit("missing").alias("summary"),
        *(
            (
                f.count(
                    f.when(
                        (f.isnan(c) | f.isnull(c)),
                        c,
                    )
                )
                / f.count("*")
                * 100
                if isinstance(t, NumericType)  # isnan only works for numeric types
                else f.count(
                    f.when(
                        f.isnull(c),
                        c,
                    )
                )
                / f.count("*")
                * 100
            )
            .cast("double")
            .alias(c)
            for c, t in df_dtypes(df)
        ),
    ).collect()
    q.put(result)

# caching the dataframe to reuse it for all the jobs?
df.cache()
# I use a queue to retrieve the results from the threads
q = Queue()
threads = [
    InheritableThread(target=fun, args=(df, q)).start()
    for fun in (get_min, get_max, get_missing, get_std)
]
# and then some code to recover the results from the queue

疑问:该方案不会生成大量带后缀的列,但不确定GIL是否会影响其并行性。


方案对比与最优选择

1. 方案一(单Select查询)的优势与疑问解答

  • 并行效率最高:Spark会将所有聚合操作合并为一个Job,仅扫描一次数据源,所有统计量在Executor端并行计算,完全贴合Spark分布式计算的设计理念,是IO效率最优的实现方式。
  • 列数量的性能影响可忽略:虽然会生成列数×4的结果列,但最终结果仅一行统计值,数据量极小,不会成为性能瓶颈。Spark处理这类聚合后的小结果集开销极低,远小于多次扫描数据源的成本。

2. 方案二(Python线程)的核心问题

  • GIL影响是次要的,多Job开销才是关键:Python线程在Driver端运行,每个线程提交独立的Spark Job。这些Job要么被Spark调度器排队执行,要么同时运行但各自独立扫描数据源(即使缓存读取也有额外开销),多次扫描的成本远高于单Job的一次扫描。
  • GIL的实际影响有限:Driver端线程在等待Spark Job结果时处于阻塞状态,此时GIL会释放,线程等待阶段不会有GIL问题。但线程仅用于提交Job,无法让Spark的计算真正并行——真正的并行由集群资源决定,多Job反而会带来调度和重复扫描的额外开销。

3. 最优方案与优化建议

优先选择方案一,并优化结果格式以方便转为JSON:

  • 优化结果结构:将单行列结构转为键值对格式,用struct整合每列的统计量,再直接转为JSON。
  • 示例优化代码:
from typing import List, Tuple
import json

def get_column_stats(col_name: str, col_type: DataType):
    # 为单个列生成所有统计量的struct
    stats = []
    # 缺失值占比
    if isinstance(col_type, NumericType):
        missing = (f.count(f.when(f.isnan(col_name) | f.isnull(col_name), col_name)) / f.count("*") * 100).cast("double").alias("missing_pct")
    else:
        missing = (f.count(f.when(f.isnull(col_name), col_name)) / f.count("*") * 100).cast("double").alias("missing_pct")
    stats.append(missing)
    # min
    if isinstance(col_type, (NumericType, DateType, TimestampType)):
        stats.append(f.min(col_name).cast(col_type).alias("min"))
    else:
        stats.append(f.lit(None).alias("min"))
    # max
    if isinstance(col_type, (NumericType, DateType, TimestampType)):
        stats.append(f.max(col_name).cast(col_type).alias("max"))
    else:
        stats.append(f.lit(None).alias("max"))
    # stddev
    if isinstance(col_type, NumericType):
        stats.append(f.stddev(col_name).cast(col_type).alias("stddev"))
    else:
        stats.append(f.lit(None).alias("stddev"))
    return f.struct(*stats).alias(col_name)

# 生成每个列的统计struct
stats_exprs = [get_column_stats(c, t) for c, t in df_dtypes(df)]
# 执行查询并转为JSON
result_df = df.select(f.struct(*stats_exprs).alias("column_stats"))
result_json = result_df.toJSON().first()
print(json.dumps(json.loads(result_json), indent=2))

该代码输出的JSON结构清晰,每个列的统计量整合在对应键下,方便后续使用。

额外注意点

  • 超大型DataFrame场景下,无需额外缓存(方案一仅扫描一次数据),若需重复使用统计结果可缓存最终的小结果集。
  • 若部分列有特殊统计逻辑,可单独调整对应列的表达式,整体框架保持不变。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 13:01:42