如何用PySpark SQL函数生成0/1数组列,判断元素是否存在于另一数组
解决方案
可以通过PySpark内置的数组与聚合函数实现需求,无需自定义UDF,具体步骤如下:
- 为每行生成唯一标识,用于后续聚合还原数组结构
- 将
col1拆分为单个元素的行,同时关联对应行的col2 - 判断每个元素是否存在于
col2中,生成0/1标记 - 按唯一标识聚合标记,还原为目标数组
col3
代码实现
from pyspark.sql import SparkSession from pyspark.sql.functions import ( monotonically_increasing_id, explode, array_contains, when, collect_list ) # 初始化SparkSession spark = SparkSession.builder.appName("ArrayMatch").getOrCreate() # 创建示例DataFrame data = [ ([1, 2, 3], [1, 4, 3]), ([5, 4, 3], [5]) ] df = spark.createDataFrame(data, ["col1", "col2"]) # 生成唯一行ID,用于后续聚合匹配 df_with_id = df.withColumn("row_id", monotonically_increasing_id()) # 拆分col1为单个元素,判断元素是否在col2中并生成标记 exploded_df = df_with_id.select( "row_id", explode("col1").alias("element"), "col2" ).withColumn( "flag", when(array_contains("col2", "element"), 1).otherwise(0) ) # 按行ID聚合标记,还原为目标数组col3 result_df = exploded_df.groupBy("row_id").agg( collect_list("flag").alias("col3") ).join(df_with_id, on="row_id").select("col3") # 展示结果 result_df.show(truncate=False)
输出结果
+---------+ |col3 | +---------+ |[1, 0, 1]| |[1, 0, 0]| +---------+
关键函数说明
monotonically_increasing_id():生成唯一行标识,确保拆分后的元素能准确对应回原行explode():将数组列拆分为多行,每个元素单独成一行,同时保留原数组的元素顺序array_contains():检查单个元素是否存在于目标数组中,配合when()生成0/1标记collect_list():按行ID聚合标记,还原为与col1元素顺序一致的数组
内容的提问来源于stack exchange,提问作者GadaaDhaariGeek
相关产品推荐
相关产品推荐

