You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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),不支持动态引用列值,所以可以用以下两种方式实现:

方法一:自连接(适合大数据量场景)

  1. 先给每个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))
  1. 自连接同分区的数据,筛选出目标行的行号在当前行号到当前行号+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"
)
  1. 按原行的字段分组,求和得到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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.22 15:55:36