如何在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
相关产品推荐
相关产品推荐

