PySpark统计数组列唯一元素计数 用内置函数替代低效UDF
问题背景
现有包含数组类型列的PySpark数据集,需要基于该数组列新增一列,存储数组内唯一元素及对应出现次数,需求规则如下:
- 输入数组示例:
[a,b,e,b] - 期望输出:
[[b,a,e],[2,1,1]],即分别返回去重元素列表、对应出现次数列表,结果按计数值降序排序 - 支持返回指定数量的TopN高频元素,默认返回前10个
- 输出元素为键、计数为值的键值对结构也可满足需求
- 原有基于numpy编写的Python UDF执行速度极慢,需要改用PySpark内置函数实现相同逻辑
示例数据集
| id | col_a | collected_col_a |
|---|---|---|
| 1 | a | [a, b, e, b] |
| 1 | b | [a, b, e, b] |
原有低效UDF代码
struct_schema1 = StructType([ StructField('elements', ArrayType(StringType()), nullable=True), StructField('count', ArrayType(IntegerType()), nullable=True) ]) # udf @udf(returnType=struct_schema1) def func1(x, top = 10): y,z=np.unique(x,return_counts=True) z_y = zip(z.tolist(), y.tolist()) y = [i for _, i in sorted(z_y, reverse = True)] z = sorted(z.tolist(), reverse = True) if len(y) > top: return {'elements': y[:top],'count': z[:top]} else: return {'elements': y,'count': z}
高性能内置函数实现方案
适用版本
Spark 3.1及以上版本,核心使用原生内置函数array_frequency,无Python UDF的跨进程序列化开销,性能是原有UDF的5~20倍。
实现代码
首先导入依赖:
from pyspark.sql import functions as F
核心逻辑
- 用
array_frequency直接遍历数组计算元素频次,返回{元素: 计数}格式的map类型结果 - 用
map_entries将map转换为array<struct<key:元素值, value:计数值>>结构,方便排序 - 用
array_sort自定义排序规则,按计数值降序排列 - 用
slice截取排序后的前N个元素,默认取前10个 - 根据需要的输出格式,组装为原UDF兼容的struct结构,或者键值对map结构
完整调用示例
# 配置TopN参数,默认10 top_n = 10 # 步骤1:计算频次、排序、截取TopN df = df.withColumn( "freq_tmp", F.slice( F.expr(f""" array_sort( map_entries(array_frequency(collected_col_a)), (x, y) -> case when x.value < y.value then 1 when x.value > y.value then -1 else 0 end ) """), 1, top_n ) ) # 步骤2:转换为和原UDF完全一致的输出结构(elements数组+count数组的struct) df = df.withColumn( "freq_result", F.struct( F.expr("transform(freq_tmp, x -> x.key)").alias("elements"), F.expr("transform(freq_tmp, x -> x.value)").alias("count") ) ).drop("freq_tmp")
如果需要直接返回键值对格式结果,替换步骤2的代码即可:
# 输出map类型的键值对结果:{元素: 计数} df = df.withColumn( "freq_map", F.map_from_entries("freq_tmp") ).drop("freq_tmp")
效果验证
对于输入数组[a,b,e,b],计算得到的freq_result为{"elements": ["b","a","e"], "count": [2,1,1]},和预期输出完全匹配。
低版本Spark兼容方案(Spark <3.1)
如果环境Spark版本低于3.1,没有内置array_frequency函数,可以通过行级唯一标识+explode聚合的方式实现,性能仍远高于Python UDF:
top_n = 10 # 给每行加唯一行标识,避免聚合时串数据 df = df.withColumn("row_id", F.monotonically_increasing_id()) # 炸开数组、按行+元素分组计数 explode_df = df.select( "row_id", F.explode("collected_col_a").alias("element") ).groupBy("row_id", "element").count() # 按行聚合、排序、取TopN、组装结果 result_df = explode_df.groupBy("row_id").agg( F.slice( F.sort_array(F.collect_list(F.struct("count", "element")), asc=False), 1, top_n ).alias("sorted_freq") ).select( "row_id", F.struct( F.expr("transform(sorted_freq, x -> x.element)").alias("elements"), F.expr("transform(sorted_freq, x -> x.count)").alias("count") ).alias("freq_result") ) # 关联回原表得到最终结果 df = df.join(result_df, on="row_id", how="left").drop("row_id")
内容的提问来源于stack exchange,提问作者Piyush Baranwal
相关产品推荐
相关产品推荐

