PySpark中如何生成数组列的两两组合
PySpark实现数组列的两两组合
需求场景
需要将DataFrame中数组列的元素生成所有不重复的两两组合,示例输入输出如下:
输入DataFrame
from pyspark.sql import SparkSession spark = SparkSession.builder.appName("array-combinations").getOrCreate() df = spark.createDataFrame( [([0, 1],), ([2, 3, 4],), ([5, 6, 7, 8],)], ['array_col'])
预期输出
+------------+------------------------------------------------+ |array_col |out | +------------+------------------------------------------------+ |[0, 1] |[[0, 1]] | |[2, 3, 4] |[[2, 3], [2, 4], [3, 4]] | |[5, 6, 7, 8]|[[5, 6], [5, 7], [5, 8], [6, 7], [6, 8], [7, 8]]| +------------+------------------------------------------------+
解决方案
方法1:利用Spark内置高阶函数(推荐)
Spark的数组列底层基于Scala的Seq,可以直接通过expr调用Scala原生的combinations(2)方法,该方法会返回数组中所有长度为2的不重复组合(元素顺序保持原数组中的相对顺序,无重复组合)。
代码实现:
from pyspark.sql.functions import expr result_df = df.withColumn("out", expr("array_col.combinations(2)")) result_df.show(truncate=False)
方法2:使用Python UDF(兼容低版本Spark)
如果你的Spark版本不支持直接调用Scala数组方法,可以用Python的itertools.combinations实现UDF:
from itertools import combinations from pyspark.sql.functions import udf from pyspark.sql.types import ArrayType, IntegerType # 定义UDF:输入数组,返回两两组合的数组 comb_udf = udf(lambda arr: list(combinations(arr, 2)), ArrayType(ArrayType(IntegerType()))) result_df = df.withColumn("out", comb_udf(df.array_col)) result_df.show(truncate=False)
注意:UDF的性能低于Spark内置函数,数据量较大时优先选择方法1。
内容的提问来源于stack exchange,提问作者ZygD
相关产品推荐
相关产品推荐

