基于布尔值转换PySpark DataFrame的实现方案咨询
PySpark 将布尔列转换为按值分组的列名集合
问题重现
给定如下PySpark DataFrame:
from pyspark.sql import Row from pyspark.sql import SparkSession spark = SparkSession.builder.appName("BooleanColumns").getOrCreate() df = spark.createDataFrame([ Row(a=True, b=False, c=False, d=True, e=False, id=1, dummy='dummy'), Row(a=False, b=False, c=True, d=True, e=False, id=2, dummy='dummy') ]) df.show()
输出:
+-----+-----+-----+----+-----+---+-----+ | a| b| c| d| e| id|dummy| +-----+-----+-----+----+-----+---+-----+ | true|false|false|true|false| 1|dummy| |false|false| true|true|false| 2|dummy| +-----+-----+-----+----+-----+---+-----+
需要转换为以下两种目标格式:
- 列表形式:
[Row(id=1, true=['a', 'd'], false=['b','c','e']), Row(id=2, true=['c','d'], false=['a','b','e'])]
- 逗号分隔字符串形式:
+---+-----+----+ |id |false|true| +---+-----+----+ |1 |b,c,e|a,d | |2 |a,b,e|c,d | +---+-----+----+
解决方案
通过宽表转长表 + 分组聚合 + 透视的组合操作实现,步骤如下:
步骤1:筛选目标布尔列
先过滤掉非布尔类型的列(如id、dummy),仅保留需要处理的布尔列:
# 获取所有布尔类型的列名 bool_cols = [col for col, dtype in df.dtypes if dtype == 'boolean'] # 保留id和布尔列,剔除dummy列 df_filtered = df.select("id", *bool_cols)
步骤2:宽表转长表(stack函数)
使用stack函数将每个布尔列拆分为列名和布尔值的键值对行:
from pyspark.sql.functions import expr # 构造stack表达式:stack(列数, 列1名, 列1值, 列2名, 列2值...) stack_expr = f"stack({len(bool_cols)}, " + ", ".join([f"'{c}', {c}" for c in bool_cols]) + ")" df_long = df_filtered.select("id", expr(stack_expr).alias("col_name", "col_value"))
转换后的长表结构:
+---+--------+---------+ |id |col_name|col_value| +---+--------+---------+ |1 |a |true | |1 |b |false | |1 |c |false | |1 |d |true | |1 |e |false | |2 |a |false | |2 |b |false | |2 |c |true | |2 |d |true | |2 |e |false | +---+--------+---------+
步骤3:分组聚合 + 透视得到结果
按id和布尔值分组,收集对应列名,再通过pivot转回宽表:
from pyspark.sql.functions import collect_list, concat_ws # 生成列表形式的结果 df_result_list = df_long.groupBy("id", "col_value") \ .agg(collect_list("col_name").alias("cols")) \ .groupBy("id") \ .pivot("col_value") \ .agg(expr("first(cols)")) # 生成逗号分隔字符串形式的结果 df_result_str = df_long.groupBy("id", "col_value") \ .agg(concat_ws(",", collect_list("col_name")).alias("cols")) \ .groupBy("id") \ .pivot("col_value") \ .agg(expr("first(cols)"))
查看最终输出
# 列表形式结果 df_result_list.show(truncate=False)
输出:
+---+---------+-------+ |id |false |true | +---+---------+-------+ |1 |[b, c, e]|[a, d] | |2 |[a, b, e]|[c, d] | +---+---------+-------+
# 逗号分隔字符串形式结果 df_result_str.show(truncate=False)
输出:
+---+-----+----+ |id |false|true| +---+-----+----+ |1 |b,c,e|a,d | |2 |a,b,e|c,d | +---+-----+----+
关键说明
stack是实现宽表转长表的核心,能高效将多列转换为键值对行collect_list用于收集同组列名,若需要字符串格式则用concat_ws拼接pivot负责将布尔值转换为列名,还原为目标宽表结构
内容的提问来源于stack exchange,提问作者RAVITEJA SATYAVADA
相关产品推荐
相关产品推荐

