PySpark如何高效为数组列的全局唯一元素分配整数编码
PySpark数组列全局唯一值编码最优实现
你之前用RDD维护全局集合逐行遍历的方式性能差,核心原因是Python RDD的序列化开销、逐行迭代的执行效率远低于Spark原生的向量化算子,用下面的方案可以把性能提升一个数量级以上:
核心思路
不需要逐行维护全局集合,只需要做一次轻量的打平去重拿到全局唯一属性的有序列表,再用Spark SQL内置的数组高阶函数批量完成映射替换,全程走Catalyst优化的原生执行路径。
实现代码
from pyspark.sql import SparkSession from pyspark.sql.functions import explode, collect_list, array_position, col, lit, transform spark = SparkSession.builder.getOrCreate() # 1. 构造原始测试数据集 raw_data = [ (101, ['a','b','c']), (102, ['a','b','d']), (103, ['b','c']), (104, ['c','e','f']), (105, ['a','b','c']), (106, ['c','g','h']), (107, ['b','d']), (108, ['d','g','i']) ] df = spark.createDataFrame(raw_data, schema=["idx", "attributes"]) # 2. 生成全局有序属性列表:打平数组、去重、按字母排序,正好对应a=0、b=1的编号规则 global_attrs = df.select(explode("attributes").alias("attr"))\ .distinct()\ .orderBy("attr")\ .agg(collect_list("attr"))\ .first()[0] # 3. 批量转换数组列:用transform高阶函数对数组内每个元素取对应下标 # 注:array_position的返回值从1开始计数,减1后得到从0开始的目标编号 result_df = df.withColumn( "attributes", transform(col("attributes"), lambda item: array_position(lit(global_attrs), item) - 1) )
执行result_df.show(truncate=False)即可得到你需要的输出:
+---+---------+ |idx|attributes| +---+---------+ |101|[0, 1, 2]| |102|[0, 1, 3]| |103|[1, 2] | |104|[2, 4, 5]| |105|[0, 1, 2]| |106|[2, 6, 7]| |107|[1, 3] | |108|[3, 6, 8]| +---+---------+
性能优化点
- 如果全局唯一属性的量级较大,可以把
global_attrs包装成广播变量,避免每个Task重复拷贝数据,进一步降低内存开销:broadcast_attrs = spark.sparkContext.broadcast(global_attrs) # 转换时取broadcast_attrs.value传入即可 - 整个流程只有一次打平去重的Shuffle操作,数组转换逻辑全部由Spark原生算子完成,没有Python层面逐行遍历的序列化、迭代开销,在大数据量下性能远高于RDD实现。
- 如果不需要严格按字母顺序分配编号,还可以直接用MLlib的
StringIndexer组件完成数组编码,代码更简洁,只需要注意设置stringOrderType="alphabetAsc"即可匹配你当前的编号规则。
内容的提问来源于stack exchange,提问作者helloworld
相关产品推荐
相关产品推荐

