如何基于多列DataFrame生成列联表(Contingency Table)
多列生成列联表的Spark实现
原始数据
原始DataFrame结构如下:
| Id | A | B | C |
|---|---|---|---|
| 1 | 0 | 0 | 1 |
| 2 | 1 | 0 | 0 |
| 3 | 1 | 0 | 1 |
| 4 | 0 | 0 | 1 |
目标列联表
需要生成的统计格式:
| Label | T/F | count |
|---|---|---|
| A | 0 | 2 |
| A | 1 | 2 |
| B | 0 | 4 |
| B | 1 | 0 |
| C | 0 | 1 |
| C | 1 | 3 |
解决方案
你已经掌握了单列分组统计,要实现多列统计,核心是先把宽表转换为长表,再分组计数,具体实现如下:
代码实现
from pyspark.sql import functions as F # 1. 宽表转长表:stack(3, ...) 中的3是要转换的列总数 long_df = data_frame.select( F.expr("stack(3, 'A', A, 'B', B, 'C', C) as (Label, `T/F`)") ) # 2. 按Label和T/F分组统计数量 result_df = long_df.groupBy("Label", "T/F").count() # 3. 补全缺失的分组(比如B列的1值,确保所有Label与0/1的组合都存在) all_combinations = spark.createDataFrame( [("A", 0), ("A", 1), ("B", 0), ("B", 1), ("C", 0), ("C", 1)], ["Label", "T/F"] ) final_result = all_combinations.join(result_df, on=["Label", "T/F"], how="left") \ .fillna(0, subset=["count"]) \ .orderBy("Label", "T/F") final_result.show()
代码说明
stack(n, col1_name, col1_value, col2_name, col2_value...):n代表要转换的列数,每一组参数对应列名和列值,将宽表的多列“堆叠”为长表的两行数据。- 分组统计后,部分组合可能因原始数据无对应值而缺失(比如B列的1),通过生成所有可能的组合再左连接,并用
fillna补0,就能得到完整的列联表。
内容的提问来源于stack exchange,提问作者Jessie
相关产品推荐
相关产品推荐

