如何在PySpark中实现类似Pandas的分组移位滚动窗口均值计算?
将Pandas滚动均值逻辑迁移到PySpark DataFrame
需求说明
需要实现和以下Pandas代码完全一致的逻辑:按day和Article分组,对QTY字段先执行1位移位,再计算窗口大小为10的滚动均值。
原Pandas代码:
df.groupby(["day","Article"])['QTY'].transform( lambda x: x.shift(1).rolling(window=10).mean())
PySpark实现步骤
PySpark依赖窗口函数(Window API)实现这类分组滚动计算,核心是定义正确的窗口范围,同时必须明确分组内的排序字段(Pandas默认按索引排序,PySpark无默认排序,需手动指定)。
- 导入必要函数
from pyspark.sql import Window from pyspark.sql.functions import avg, count, when
- 定义窗口规范
替换示例中的sort_col为你数据实际的排序字段(比如交易日期、序列ID等),确保分组内的行顺序和Pandas逻辑一致:
# 分区:按day和Article分组;排序:按业务逻辑指定字段;窗口范围:取当前行的前10行到前1行(共10行) window_spec = Window.partitionBy("day", "Article")\ .orderBy("sort_col")\ .rowsBetween(-10, -1)
- 计算滚动均值(匹配Pandas默认行为)
Pandas默认要求窗口内有10个非缺失值才返回均值,否则返回NaN。以下代码完全匹配该逻辑:
df = df.withColumn( "rolling_avg_qty", # 当窗口内非缺失QTY的数量等于10时计算均值,否则返回None(对应Pandas的NaN) when(count("QTY").over(window_spec) == 10, avg("QTY").over(window_spec)) .otherwise(None) )
简化版本(按需选择)
如果不需要严格匹配Pandas的缺失值规则,允许不足10行时计算现有行的均值,可直接简化为:
df = df.withColumn( "rolling_avg_qty", avg("QTY").over(window_spec) )
内容的提问来源于stack exchange,提问作者Abad
相关产品推荐
相关产品推荐

