如何使用PySpark按分组迭代并生成符合求和条件的数组列
PySpark 实现按分组筛选符合条件的B%值汇总数组
实现思路
- 用窗口函数按
Group字段分区,先收集同组所有B%值为全局数组 - 通过高阶函数
filter遍历全局数组,筛选满足当前行A% + 数组中B值 ≤ Target%的元素,生成最终的SumArray列
代码实现
依赖导入
from pyspark.sql import SparkSession from pyspark.sql import functions as F from pyspark.sql.window import Window
测试数据构造(实际使用时替换为自有数据源即可)
spark = SparkSession.builder.appName("GroupFilterDemo").getOrCreate() data = [ ("A", 0.05, 0.85, 1.0), ("A", 0.07, 0.75, 1.0), ("A", 0.08, 0.95, 1.0), ("B", 0.03, 0.80, 1.0), ("B", 0.05, 0.83, 1.0), ("B", 0.04, 0.85, 1.0) ] df = spark.createDataFrame(data, schema=["Group", "A %", "B %", "Target %"])
核心逻辑
# 定义按Group分区的窗口 group_window = Window.partitionBy("Group") result_df = df.withColumn("group_all_b", F.collect_list("B %").over(group_window)) \ .withColumn("SumArray", F.expr("filter(group_all_b, b -> (`A %` + b) <= `Target %`)")) \ .drop("group_all_b") # 输出验证结果 result_df.show(truncate=False)
注意事项
- 由于字段名包含空格,在SQL表达式中需要用反引号包裹字段名,避免出现语法错误
- 该实现基于PySpark 2.4及以上版本,低版本需先升级PySpark环境
内容的提问来源于stack exchange,提问作者Alex Triece
相关产品推荐
相关产品推荐

