You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.19 07:45:26