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

如何在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无默认排序,需手动指定)。

  1. 导入必要函数
from pyspark.sql import Window
from pyspark.sql.functions import avg, count, when
  1. 定义窗口规范
    替换示例中的sort_col为你数据实际的排序字段(比如交易日期、序列ID等),确保分组内的行顺序和Pandas逻辑一致:
# 分区:按day和Article分组;排序:按业务逻辑指定字段;窗口范围:取当前行的前10行到前1行(共10行)
window_spec = Window.partitionBy("day", "Article")\
                    .orderBy("sort_col")\
                    .rowsBetween(-10, -1)
  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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 03:25:23