PySpark如何提取数组中单个或求和等于指定目标值的元素
PySpark 数组列筛选符合条件元素解决方案
现有PySpark DataFrame 包含存储double类型元素的数组列,以及目标值列,需要为每行生成新数组,收录所有满足「单个元素等于目标值」或「多个元素求和等于目标值」的元素,效果和需求示例一致。
注意事项
由于double类型存在浮点精度误差,匹配时需设置容差避免判断错误,容差大小可根据业务精度要求调整。
实现步骤
- 导入依赖函数
from pyspark.sql import SparkSession from pyspark.sql.functions import udf from pyspark.sql.types import ArrayType, DoubleType
- 实现核心筛选UDF
# 浮点比较容差,可自行调整 TOLERANCE = 1e-6 def filter_valid_elements(arr, target): arr_length = len(arr) valid_indexes = set() # 枚举所有非空子集 for mask in range(1, 1 << arr_length): current_sum = 0.0 current_indexes = [] for i in range(arr_length): if mask & (1 << i): current_sum += arr[i] current_indexes.append(i) # 子集和匹配目标值 if abs(current_sum - target) < TOLERANCE: for idx in current_indexes: valid_indexes.add(idx) # 按原数组顺序返回结果 return [arr[i] for i in sorted(valid_indexes)] # 注册UDF filter_udf = udf(filter_valid_elements, ArrayType(DoubleType()))
- 调用UDF生成结果列
# 构造示例DataFrame spark = SparkSession.builder.appName("array_filter").getOrCreate() data = [ ([0.0001,2.5,3.0,0.0031], 0.0032), ([2.5,1.0,0.5,3.0], 3.0), ([1.0,1.0,1.5,1.0], 4.5) ] df = spark.createDataFrame(data, schema=["Array", "Target"]) # 生成NewArray列 df = df.withColumn("NewArray", filter_udf("Array", "Target")) df.show(truncate=False)
补充说明
- 该方案适合数组长度不超过20的场景,若数组过长,子集枚举的时间复杂度会指数上升,建议改用动态规划方案优化
- 运行结果和需求给出的示例完全一致
内容的提问来源于stack exchange,提问作者Alex Triece
相关产品推荐
相关产品推荐

