PySpark中如何用avg_lt列值替代rowsBetween的固定数值
PySpark动态窗口滚动求和:基于每行avg_lt列计算范围总和
原代码通过固定窗口范围(当前行及后续3行)计算expected_demand的总和,生成lt_demand列:
spec = Window. \ partitionBy("material_id"). \ orderBy("date"). \ rowsBetween(Window.currentRow, 3) rop.withColumn("lt_demand", f.sum("expected_demand").over(spec)).show()
现在需要将窗口结束的固定值3替换为每行的avg_lt列值,即对每行计算当前行及后续N行(N为该行avg_lt的值)的expected_demand之和。
因为PySpark原生的rowsBetween只支持固定数值或特殊边界(如Window.unboundedFollowing),不支持动态引用列值,所以可以用以下两种方式实现:
方法一:自连接(适合大数据量场景)
- 先给每个
material_id分区内的行按date排序,添加行号方便范围筛选:
from pyspark.sql import functions as f from pyspark.sql.window import Window row_num_window = Window.partitionBy("material_id").orderBy("date") rop_with_row = rop.withColumn("row_num", f.row_number().over(row_num_window))
- 自连接同分区的数据,筛选出目标行的行号在
当前行号到当前行号+avg_lt之间的记录:
joined = rop_with_row.alias("a").join( rop_with_row.alias("b"), (f.col("a.material_id") == f.col("b.material_id")) & (f.col("b.row_num").between(f.col("a.row_num"), f.col("a.row_num") + f.col("a.avg_lt"))), how="left" )
- 按原行的字段分组,求和得到
lt_demand:
result = joined.groupBy("a.material_id", "a.date", "a.expected_demand", "a.avg_lt", "a.row_num") \ .agg(f.sum("b.expected_demand").alias("lt_demand")) \ .drop("row_num") \ .orderBy("material_id", "date") result.show()
方法二:数组聚合(数据量较小时更高效)
通过收集分区内的需求数组,再根据行号截取对应范围求和:
# 先沿用方法一中的rop_with_row(带行号的数据集) collect_window = Window.partitionBy("material_id").orderBy("date") rop_with_arrays = rop_with_row.withColumn( "demand_array", f.collect_list("expected_demand").over(collect_window) ).withColumn( "row_num_array", f.collect_list("row_num").over(collect_window) ) # 截取目标范围的数组并求和 result = rop_with_arrays.withColumn( "target_demands", f.array_slice( f.col("demand_array"), f.array_position(f.col("row_num_array"), f.col("row_num")), f.col("avg_lt") + 1 # +1是因为要包含当前行 ) ).withColumn( "lt_demand", f.expr("aggregate(target_demands, 0D, (acc, x) -> acc + x)") ).drop("row_num", "demand_array", "row_num_array", "target_demands") result.show()
内容的提问来源于stack exchange,提问作者Yasmine Elserafy
相关产品推荐
相关产品推荐

