滚动窗口计算标准差时内存占用过高问题问询
我有一个二维数组,想要对尺寸约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,故询问该结果的原因。
原因分析
标准差计算的本质复杂度:
均值计算仅需一次遍历窗口累加求和后取平均,而标准差计算依赖公式sqrt(E[X²] - (E[X])²),需要先计算整个数组的滚动均值和滚动平方均值,再推导得到标准差。这意味着标准差计算需要同时保留两组全尺寸的中间结果,内存占用至少是均值计算的两倍。懒计算的执行顺序问题:
你的代码是先计算全数组的滚动统计量,再通过isel提取子集。对于大窗口,全数组的滚动标准差会生成与原数组等大的中间结果,在提取子集前会占用大量内存;而均值计算的中间结果虽然同样大小,但单数组的内存压力远低于标准差所需的双数组+最终结果的组合。Chunk与窗口的匹配冲突:
你设置的Chunk大小为256x256,当窗口大小(如41)接近Chunk尺寸的1/6时,Dask需要加载更多相邻Chunk来完成跨边界的窗口计算。对于标准差,跨Chunk的中间数据保留量更大,进一步推高内存占用。版本优化不足:
你使用的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

