如何统计PySpark数组列中指定单词列表的总出现次数
PySpark 统计数组类型列指定单词总出现次数实现
需求说明
给定待匹配单词列表ls=['aa','bb'],DataFrame的item列为数组类型,需要新增count列,统计每行item数组中属于待匹配列表的单词总出现次数,示例预期输出如下:
| item | count |
|---|---|
| [aa, bb] | 2 |
| [aa, bc] | 1 |
| [ad, bc] | 0 |
前置准备
首先构造测试环境与示例数据:
from pyspark.sql import SparkSession from pyspark.sql.functions import col, lit, aggregate, array_intersect, size # 初始化SparkSession spark = SparkSession.builder.appName("array_count").getOrCreate() # 构造示例DataFrame data = [ (["aa", "bb"],), (["aa", "bc"],), (["ad", "bc"],) ] df = spark.createDataFrame(data, schema=["item"]) # 定义待匹配单词列表 match_words = ["aa", "bb"]
方案1:高阶函数aggregate实现(推荐)
无需打散数组,性能最优,同时支持数组内重复单词的计数:
df_res = df.withColumn( "count", aggregate( col("item"), lit(0), # 初始计数为0 lambda acc, curr: acc + curr.isin(match_words).cast("int") ) ) df_res.show()
输出结果:
+--------+-----+ | item|count| +--------+-----+ |[aa, bb]| 2| |[aa, bc]| 1| |[ad, bc]| 0| +--------+-----+
方案2:简化写法(仅适用于数组无重复待匹配单词场景)
如果确定item列数组内不会出现重复的待匹配单词,可以用array_intersect直接计算交集大小:
df_res = df.withColumn("count", size(array_intersect(col("item"), lit(match_words))))
内容的提问来源于stack exchange,提问作者Abhishek
相关产品推荐
相关产品推荐

