PySpark中基于Window函数按条件收集id2列表的实现问题
解决PySpark分组收集指定level范围id2的问题
你的核心问题是Window范围设置错误,导致收集了全组的id2值。结合你的期望结果(实际需收集同组内level大于等于当前行level的id2,你描述中的"小于等于"应为笔误),可通过两种方式实现:
方法一:调整Window的行范围(推荐,性能更优)
既然已按level升序排序,只需将Window的行范围设置为当前行到分组末尾,就能直接收集符合条件的id2:
import pyspark.sql.functions as F df = spark.createDataFrame([ ("A", 0, "M1", "D1"), ("A", 1, "D1", "D2"), ("A", 2, "D2", "D3"), ("A", 3, "D3", "D4"), ("B", 0, "M2", "D5"), ("B", 1, "D4", "D6"), ("B", 2, "D5", "D7") ], ["group_id", "level", "id1", "id2"]) # 调整Window范围:当前行到分组末尾 window = Window.partitionBy('group_id').orderBy('level').rowsBetween( Window.currentRow, Window.unboundedFollowing ) df_with_list = df.withColumn( "list_lower_level", F.collect_list("id2").over(window) ) df_with_list.show()
方法二:数组过滤(适用于level非连续的场景)
如果你的level存在跳变(比如不是连续的0、1、2...),可先收集全组的(level, id2)对,再通过数组过滤提取符合条件的id2:
# 先收集全组的level和id2的结构体数组 window_full = Window.partitionBy('group_id') df_with_full_list = df.withColumn( "full_level_id2", F.collect_list(F.struct("level", "id2")).over(window_full) ) # 过滤出level >= 当前行level的id2,再提取id2字段 df_with_list = df_with_full_list.withColumn( "list_lower_level", F.transform( F.filter("full_level_id2", lambda x: x.level >= F.col("level")), lambda x: x.id2 ) ).drop("full_level_id2") df_with_list.show()
两种方法都会得到你期望的结果:
+--------+-----+---+---+----------------+ |group_id|level|id1|id2|list_lower_level| +--------+-----+---+---+----------------+ | A| 0| M1| D1|[D1, D2, D3, D4]| | A| 1| D1| D2| [D2, D3, D4]| | A| 2| D2| D3| [D3, D4]| | A| 3| D3| D4| [D4]| | B| 0| M2| D5| [D5, D6, D7]| | B| 1| D4| D6| [D6, D7]| | B| 2| D5| D7| [D7]| +--------+-----+---+---+----------------+
内容的提问来源于stack exchange,提问作者Henri
相关产品推荐
相关产品推荐

