Spark 2.3.0下PySpark DataFrame的Array列元素如何按字符串条件筛选
PySpark 2.3.0 实现Array列元素按字符串规则过滤
需求说明
对DataFrame的Array类型列中每行的数组元素做过滤,仅保留包含apple子串、或是以app开头的元素。
前置准备:构造示例DataFrame
先给出示例数据构造代码,方便验证效果:
from pyspark.sql import SparkSession from pyspark.sql.functions import udf, col, explode, collect_list, monotonically_increasing_id from pyspark.sql.types import ArrayType, StringType spark = SparkSession.builder.appName("array_filter").getOrCreate() # 构造示例数据 data = [ (["apple", "banana", "orange"],), (["strawberry", "raspberry"],), (["apple", "pineapple", "grapes"],) ] df = spark.createDataFrame(data, ["Array Col"]) df.show(truncate=False)
实现方案
Spark 2.3.0版本尚未内置数组高阶过滤函数,可通过以下两种常用方案实现需求:
方案1:自定义UDF实现(推荐,兼容性最好)
UDF实现逻辑直观,代码量小:
# 定义数组过滤逻辑 def filter_array(arr): if not arr: return [] res = [] for elem in arr: if 'apple' in elem or elem.startswith('app'): res.append(elem) return res # 注册UDF并指定返回类型 filter_udf = udf(filter_array, ArrayType(StringType())) # 应用UDF得到过滤后的结果 df_result = df.withColumn("Filtered Array Col", filter_udf(col("Array Col"))) df_result.show(truncate=False)
方案2:无UDF实现(通过explode+聚合实现)
如果不想引入UDF,可以通过数组展开、过滤、再聚合的方式实现:
# 给每行添加唯一ID,用于后续聚合还原行 df_with_id = df.withColumn("row_id", monotonically_increasing_id()) # 把数组列展开为单个元素的行 df_exploded = df_with_id.select("row_id", explode(col("Array Col")).alias("elem")) # 过滤符合匹配规则的元素 df_filtered = df_exploded.filter(col("elem").contains("apple") | col("elem").startswith("app")) # 按行ID聚合,把符合条件的元素重新组装为数组 df_agg = df_filtered.groupBy("row_id").agg(collect_list("elem").alias("Filtered Array Col")) # 关联回原表得到完整结果 df_result = df_with_id.join(df_agg, on="row_id", how="left").drop("row_id") df_result.show(truncate=False)
预期输出
两种方案最终输出结果一致:
+----------------------------+-------------------+ |Array Col |Filtered Array Col | +----------------------------+-------------------+ |[apple, banana, orange] |[apple] | |[strawberry, raspberry] |[] | |[apple, pineapple, grapes] |[apple, pineapple] | +----------------------------+-------------------+
补充说明:如果后续升级到Spark 2.4及以上版本,可以直接使用内置高阶函数
filter实现,无需自定义UDF:from pyspark.sql.functions import expr df_result = df.withColumn("Filtered Array Col", expr("filter(`Array Col`, x -> x contains 'apple' or x starts with 'app')"))
内容的提问来源于stack exchange,提问作者AnonymousMe
相关产品推荐
相关产品推荐

