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

滚动窗口计算标准差时内存占用过高问题问询

问题:xarray滚动窗口标准差计算内存占用过高崩溃

我有一个二维数组,想要对尺寸约50像素及以上的空间窗口(win_size)计算均值和标准差,仅提取存储在两个数组中的像素子集的结果。我的代码如下:

# 初始化中心对齐的滚动窗口
R_spatial = R.rolling({"x": win_size, "y": win_size}, center=True)
# 计算指定像素的均值
R_mean = R_spatial.mean().isel(x=x_loc_idx, y=y_loc_idx).compute()
# 计算指定像素的标准差
R_std = R_spatial.std().isel(x=x_loc_idx, y=y_loc_idx).compute()

均值计算无问题,但标准差计算会持续占用内存直至Python解释器崩溃。使用大型集群可解决该问题,但我认为xarray/dask本应能处理无法装入内存的计算,请问原因是什么?使用的xarray和dask版本为2023.05。

内存分析

使用mprof run run_me.py进行基准测试,run_me.py代码如下:

import numpy as np
import xarray as xr
from memory_profiler import profile


def create_data(N = 3500, n_samps = 5):
    R = xr.DataArray(np.random.randn(N, N), dims=["x", "y"]).chunk({"x":256, "y":256})
    x_loc_idx, y_loc_idx = np.random.randint(0,N, (2, n_samps))
    return R, x_loc_idx, y_loc_idx

@profile
def do_mean(R, x_loc_idx, y_loc_idx, win_size = 21):
    # 初始化中心对齐的滚动窗口
    R_spatial = R.rolling({"x": win_size, "y": win_size}, center=True)
    # 计算指定像素的均值
    return R_spatial.mean().isel(x=x_loc_idx, y=y_loc_idx).compute()

@profile
def do_std(R, x_loc_idx, y_loc_idx, win_size = 21):
    # 初始化中心对齐的滚动窗口
    R_spatial = R.rolling({"x": win_size, "y": win_size}, center=True)
    # 计算指定像素的标准差
    return R_spatial.std().isel(x=x_loc_idx, y=y_loc_idx).compute()

if __name__ == "__main__":
    R, x_loc_idx, y_loc_idx = create_data()
    for win_size in [11, 31]:
        mu = do_mean(R, x_loc_idx, y_loc_idx, win_size=win_size)
        sigma = do_std(R, x_loc_idx, y_loc_idx, win_size=win_size)

在我的系统上运行结果如下:

$ mprof run run_me.py 
mprof: Sampling memory every 0.1s
running new process
running as a Python program...
Filename: run_me.py

Line #    Mem usage    Increment  Occurrences   Line Contents
=============================================================
    11    295.4 MiB    295.4 MiB           1   @profile
    12                                         def do_mean(R, x_loc_idx, y_loc_idx, win_size = 21):
    13                                             # 初始化中心对齐的滚动窗口
    14    295.5 MiB      0.1 MiB           1       R_spatial = R.rolling({"x": win_size, "y": win_size}, center=True)
    15                                             # 计算指定像素的均值
    16    344.0 MiB     48.5 MiB           1       return R_spatial.mean().isel(x=x_loc_idx, y=y_loc_idx).compute()


Filename: run_me.py

Line #    Mem usage    Increment  Occurrences   Line Contents
=============================================================
    18    344.0 MiB    344.0 MiB           1   @profile
    19                                         def do_std(R, x_loc_idx, y_loc_idx, win_size = 21):
    20                                             # 初始化中心对齐的滚动窗口
    21    344.0 MiB      0.0 MiB           1       R_spatial = R.rolling({"x": win_size, "y": win_size}, center=True)
    22                                             # 计算指定像素的标准差
    23    333.3 MiB    -10.8 MiB           1       return R_spatial.std().isel(x=x_loc_idx, y=y_loc_idx).compute()


Filename: run_me.py

Line #    Mem usage    Increment  Occurrences   Line Contents
=============================================================
    11    333.3 MiB    333.3 MiB           1   @profile
    12                                         def do_mean(R, x_loc_idx, y_loc_idx, win_size = 21):
    13                                             # 初始化中心对齐的滚动窗口
    14    333.3 MiB      0.0 MiB           1       R_spatial = R.rolling({"x": win_size, "y": win_size}, center=True)
    15                                             # 计算指定像素的均值
    16    345.4 MiB     12.1 MiB           1       return R_spatial.mean().isel(x=x_loc_idx, y=y_loc_idx).compute()


Filename: run_me.py

Line #    Mem usage    Increment  Occurrences   Line Contents
=============================================================
    18    345.4 MiB    345.4 MiB           1   @profile
    19                                         def do_std(R, x_loc_idx, y_loc_idx, win_size = 21):
    20                                             # 初始化中心对齐的滚动窗口
    21    345.4 MiB      0.0 MiB           1       R_spatial = R.rolling({"x": win_size, "y": win_size}, center=True)
    22                                             # 计算指定像素的标准差
    23    440.0 MiB     94.6 MiB           1       return R_spatial.std().isel(x=x_loc_idx, y=y_loc_idx).compute()

当我将win_size设为41或样本数增加到10时,程序会耗尽内存崩溃(运行在Linux笔记本上);使用Dask GatewayServers可处理预期负载(n_samps达数万,win_size约61)。该计算本很简单,我可尝试其他方案(如用numba循环扩展邻域计算或使用map_blocks等),但我想熟悉xarray,故询问该结果的原因。


原因分析

  1. 标准差计算的本质复杂度:
    均值计算仅需一次遍历窗口累加求和后取平均,而标准差计算依赖公式sqrt(E[X²] - (E[X])²),需要先计算整个数组的滚动均值和滚动平方均值,再推导得到标准差。这意味着标准差计算需要同时保留两组全尺寸的中间结果,内存占用至少是均值计算的两倍。

  2. 懒计算的执行顺序问题:
    你的代码是先计算全数组的滚动统计量,再通过isel提取子集。对于大窗口,全数组的滚动标准差会生成与原数组等大的中间结果,在提取子集前会占用大量内存;而均值计算的中间结果虽然同样大小,但单数组的内存压力远低于标准差所需的双数组+最终结果的组合。

  3. Chunk与窗口的匹配冲突:
    你设置的Chunk大小为256x256,当窗口大小(如41)接近Chunk尺寸的1/6时,Dask需要加载更多相邻Chunk来完成跨边界的窗口计算。对于标准差,跨Chunk的中间数据保留量更大,进一步推高内存占用。

  4. 版本优化不足:
    你使用的2023.05版本xarray/dask,在滚动窗口的懒计算优化上存在局限,后续版本针对滚动统计量的任务图做了改进,能更高效地处理子集提取,避免生成全数组的中间结果。

临时解决方案

  • 先提取子集邻域再计算:跳过全数组滚动统计,直接提取目标像素的窗口邻域后计算标准差,避免生成全尺寸中间结果:
    def do_std_optimized(R, x_loc_idx, y_loc_idx, win_size=21):
        half_win = win_size // 2
        stds = []
        for x, y in zip(x_loc_idx, y_loc_idx):
            # 提取目标像素的窗口邻域,处理边界
            x_slice = slice(max(0, x-half_win), min(R.shape[0], x+half_win+1))
            y_slice = slice(max(0, y-half_win), min(R.shape[1], y+half_win+1))
            window = R.isel(x=x_slice, y=y_slice).compute()
            stds.append(window.std().item())
        return np.array(stds)
    
  • 调整Chunk大小:将Chunk尺寸设为远大于窗口大小(如512x512或1024x1024),减少跨Chunk计算的开销,降低内存中同时加载的Chunk数量。
  • 升级依赖版本:更新xarray和dask到2023.10及以后版本,利用滚动窗口的优化特性。

内容的提问来源于stack exchange,提问作者Jose

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 12:52:02