如何在PySpark DataFrame中拆分数组元素并保留其余元素
PySpark实现数组元素拆分并生成排除自身的数组列
方法一:使用Spark内置函数(推荐,性能更优)
Spark 2.4及以上版本提供了array_remove函数,配合explode拆分数组就能快速完成需求:
- 先构造测试DataFrame(模拟你的输入数据):
from pyspark.sql import SparkSession from pyspark.sql.functions import explode, array_remove spark = SparkSession.builder.appName("array_transform").getOrCreate() # 输入DataFrame input_df = spark.createDataFrame([([101826, 101827, 101576],)], ["ids"]) input_df.show(truncate=False)
- 执行转换逻辑:
# 拆分数组得到单个id,同时生成排除当前id的新数组列 output_df = input_df.withColumn("id", explode("ids")) \ .withColumn("ids", array_remove("ids", "id")) \ .select("id", "ids") output_df.show(truncate=False)
方法二:使用自定义UDF
如果你的Spark版本低于2.4,或者需要更灵活的逻辑,可以用UDF实现:
from pyspark.sql.functions import udf, explode from pyspark.sql.types import ArrayType, IntegerType # 定义移除当前元素的逻辑 def exclude_self(arr, current_id): return [num for num in arr if num != current_id] # 注册UDF并指定返回类型 exclude_self_udf = udf(exclude_self, ArrayType(IntegerType())) # 执行转换 output_df = input_df.withColumn("id", explode("ids")) \ .withColumn("ids", exclude_self_udf("ids", "id")) \ .select("id", "ids") output_df.show(truncate=False)
说明
- 优先用内置函数方案,Spark对内置函数做了分布式性能优化,比自定义UDF效率更高。
- 两种方法执行后都会生成你需要的输出格式。
内容的提问来源于stack exchange,提问作者littlely
相关产品推荐
相关产品推荐

