PySpark:用单个查询统计DataFrame各列词频
用单个PySpark查询统计所有列的词频
要实现单个查询统计DataFrame所有列的词频,核心思路是将宽表转换为长表(把每列的列名和对应值映射为两列:column和word),再按这两列分组统计次数。以下是两种适配不同Spark版本的实现方案:
样例输入与预期输出
样例输入DataFrame
| col1 | col2 |
|---|---|
| apple | cat |
| banana | dog |
| apple | cat |
| orange | dog |
预期输出
| column | word | count |
|---|---|---|
| col1 | apple | 2 |
| col1 | banana | 1 |
| col1 | orange | 1 |
| col2 | cat | 2 |
| col2 | dog | 2 |
方案1:Spark 3.1+ 用stack函数(推荐)
Spark 3.1及以上版本支持stack函数,可以直接将多列转换为键值对形式,无需循环:
from pyspark.sql import SparkSession from pyspark.sql.functions import expr # 初始化SparkSession spark = SparkSession.builder.appName("AllColumnsWordCount").getOrCreate() # 构造样例DataFrame data = [("apple", "cat"), ("banana", "dog"), ("apple", "cat"), ("orange", "dog")] df = spark.createDataFrame(data, ["col1", "col2"]) # 获取所有列名 cols = df.columns # 构造stack表达式:stack(列数, '列名1', 列名1, '列名2', 列名2, ...) stack_expr = f"stack({len(cols)}, {', '.join([f'{repr(col)}, {col}' for col in cols])}) as (column, word)" # 转换为长表并统计词频 result_df = df.selectExpr(stack_expr) \ .filter("word IS NOT NULL") # 过滤空值,避免统计无效内容 .groupBy("column", "word") \ .count() \ .orderBy("column", "count", ascending=False) # 可选:按列名和次数排序 # 展示结果 result_df.show()
代码说明
stack函数将每一列的列名字符串和列值映射为column和word两列,把宽表转为长表- 过滤
null值避免统计空内容 - 按
column和word分组,调用count()统计出现次数
方案2:Spark 3.1以下版本 用unionAll合并列
如果你的Spark版本低于3.1,没有stack函数,可以通过循环遍历每一列,生成对应列的键值对DataFrame,再用unionAll合并后统计:
from pyspark.sql import SparkSession from pyspark.sql.functions import lit, col spark = SparkSession.builder.appName("AllColumnsWordCount").getOrCreate() # 构造样例DataFrame data = [("apple", "cat"), ("banana", "dog"), ("apple", "cat"), ("orange", "dog")] df = spark.createDataFrame(data, ["col1", "col2"]) # 初始化空的长表DataFrame long_df = None for col_name in df.columns: # 为当前列生成(column, word)格式的临时DataFrame temp_df = df.select( lit(col_name).alias("column"), # 列名作为column列的值 col(col_name).alias("word") # 当前列的值作为word列的值 ) # 合并到长表 if long_df is None: long_df = temp_df else: long_df = long_df.unionAll(temp_df) # 统计词频 result_df = long_df.filter("word IS NOT NULL") \ .groupBy("column", "word") \ .count() \ .orderBy("column", "count", ascending=False) result_df.show()
代码说明
- 遍历每一列,用
lit(col_name)固定列名字段,col(col_name)取列值,生成临时DataFrame - 用
unionAll将所有临时DataFrame合并为长表 - 后续过滤空值、分组统计的逻辑和方案1一致
内容的提问来源于stack exchange,提问作者venus
相关产品推荐
相关产品推荐

