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

如何在PySpark中高效生成带维度总计的多维交叉表?

摘要:有没有更优的实现方式?
columns = ['sex', 'class', 'survived'] # 适用于多列场景
grouped_crosstab = sdf.groupBy(*columns).count()

for column in columns:
    grouped_crosstab = grouped_crosstab.join(
        grouped_crosstab.groupBy(column).agg(F.sum('count').alias(f'{column}_total')),
        column,
        'left')

问题背景

在PySpark中,你可以对DataFrame使用crosstab方法生成二维交叉表。而groupBy方法则能生成多维“交叉表”,不过输出是长表(高瘦型)格式。

示例代码如下:

columns = ['x', 'y', 'z'] # 假设这些列是低基数的分类变量,而非连续值
two_dimensional_crosstab = df.crosstab(columns[0], columns[1]) # 仅对比x和y维度
multi_dimensional_view = df.groupBy(*columns).count() # 同时对比x、y、z三个维度

用示例数据可视化

import seaborn
df = seaborn.load_dataset('titanic')
sdf = spark.createDataFrame(df) # Spark上下文的配置不在本问题讨论范围内

使用的是泰坦尼克号数据集,包含性别、舱位、生存状态等分类字段。

我们基于sex和class字段分别用crosstab和groupBy生成二维交叉表:

two_d_crosstab = sdf.crosstab('sex', 'class')
grouped_crosstab = sdf.groupBy('sex', 'class').count()

两者输出格式不同:two_d_crosstab是宽表,每行对应一个性别,列是不同舱位的计数;grouped_crosstab是长表,每行是性别+舱位的组合,对应该组合的计数。

和crosstab不同,groupBy方法可以轻松扩展到多列,但需要注意其长表格式的特点。

行列总计

出于统计需求(比如调查校准),通常需要给交叉表添加行和列的总计。在二维场景下,可以通过以下方式实现:

index_column = two_d_crosstab.columns[0]
col_list = two_d_crosstab.columns[1:]
two_d_crosstab = two_d_crosstab.withColumn('column_total', sum([F.col(c) for c in col_list]))
transposed_df = two_d_crosstab.pandas_api()\
    .set_index(index_column)\
    .T.reset_index()\
    .rename(columns = {'index':index_column})\
    .to_spark()
col_list = transposed_df.columns[1:]
two_d_crosstab = transposed_df.withColumn('row_total', sum([F.col(c) for c in col_list]))

处理后的two_d_crosstab包含各列的总计值和各行的总计值,能直观展示不同维度的汇总数据。

多维总计

那如何在多维交叉表中添加各类总计呢?

我尝试了以下方法:

sex_tot = grouped_crosstab.groupBy('sex').agg(F.sum('count').alias('sex_total'))
class_tot = grouped_crosstab.groupBy('class').agg(F.sum('count').alias('class_total'))
grouped_crosstab = grouped_crosstab.join(sex_tot, 'sex', 'left').join(class_tot, 'class', 'left')

输出的长表中会包含每个性别对应的总计数,以及每个舱位对应的总计数。

当添加survived作为第三维度时:

columns = ['sex', 'class', 'survived']
grouped_crosstab = sdf.groupBy(*columns).count()

for column in columns:
    grouped_crosstab = grouped_crosstab.join(
        grouped_crosstab.groupBy(column).agg(F.sum('count').alias(f'{column}_total')),
        column,
        'left')

输出结果中会包含每个维度的总计数,但存在大量重复数据。而且随着维度列数的增加,需要执行的group by和join操作数量也会同步增加,在包含数百万行的大型DataFrame上,这种方式会非常繁琐且低效。

有没有更好(更具可扩展性)的方法?


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 07:05:35