如何基于多数投票实现xarray.coarsen降采样?
问题描述
我需要将大型GeoTIFF文件(存储树种分类的uint8整数值,形状(86145, 64074),约5.5GB,仅占内存20%)重采样至更低分辨率,要求对每个3x3或2x2窗口执行多数投票重采样(忽略NaN/填充值)。
最初用scipy.stats.mode实现,但无论用原生numpy还是Dask都会导致内存溢出。Dask无法延迟计算,会立即执行并持续分配内存直至耗尽RAM和SWAP,导致系统崩溃。
我理解的任务逻辑很清晰:
- 将数组重塑为带降采样窗口维度的形式(类似
xr.coarsen或skimage.measure.block_reduce的逻辑) - 对每个窗口的数值计算众数(多数投票)
- 生成低分辨率数据集
但尝试了np.bincount、np.unique、结合numba自行实现、二维重塑等方法都无效,想知道如何用numpy高效解决这个问题。
示例复现代码
import xarray as xr import dask.array as da import numpy as np from scipy import stats # 定义维度 nx, ny, nt = 3000, 300, 100 # 各维度大小 chunks = (300, 30, 10) # Dask分块大小 # 创建Dask数组 data1 = da.random.random((nx, ny, nt), chunks=chunks) data2 = da.random.random((nx, ny, nt), chunks=chunks) # 创建坐标 x = np.linspace(0, 10, nx) y = np.linspace(0, 5, ny) time = np.arange(nt) # 构建xarray数据集 ds = xr.Dataset( { "temperature": (("x", "y", "time"), data1), "precipitation": (("x", "y", "time"), data2), }, coords={ "x": x, "y": y, "time": time, } ) # 自定义众数计算函数 def find_mode(arr, axis): m, _ = stats.mode(arr, axis=axis) return m # 这个是延迟计算的! coarse_mean_ds = ds.coarsen(x=3, y=3, boundary='pad').reduce(np.mean) # 这个会立即计算,单个worker占用过多内存! maj_vote_coarse = ds.coarsen(x=3, y=3, boundary='pad').reduce(find_mode)
解决方案
核心思路:数组重塑+向量化np.bincount高效实现多数投票
scipy.stats.mode在处理大窗口时效率低且内存占用高,我们可以通过数组重塑拆分窗口维度,结合np.bincount的向量化统计来替代,同时针对性处理NaN/填充值。
1. 单块数组的3x3多数投票实现(支持NaN过滤)
import numpy as np def majority_vote_3x3(arr, pad_val=np.nan): H, W = arr.shape # 补全边界至窗口大小的整数倍 pad_h = (3 - H % 3) % 3 pad_w = (3 - W % 3) % 3 arr_padded = np.pad(arr, ((0, pad_h), (0, pad_w)), mode='constant', constant_values=pad_val) # 重塑为窗口维度:(H//3, W//3, 9),每个元素对应一个3x3窗口的所有值 arr_reshaped = arr_padded.reshape(H//3, 3, W//3, 3).transpose(0,2,1,3).reshape(-1, 9) # 对每个窗口统计非填充值的众数 def get_mode(window): # 过滤填充值 valid_vals = window[~np.isnan(window)] if np.isnan(pad_val) else window[window != pad_val] if len(valid_vals) == 0: return pad_val # uint8值范围0-255,用bincount统计效率远高于mode return np.argmax(np.bincount(valid_vals, minlength=256)) mode_vals = np.apply_along_axis(get_mode, axis=1, arr=arr_reshaped) # 重塑回降采样后的形状 return mode_vals.reshape(H//3, W//3)
2. 大数组分块处理(控制内存峰值)
即使数据仅占内存20%,直接重塑大数组仍可能触发临时内存占用过高,建议分块处理:
def majority_vote_large_array(arr, window_size=3, pad_val=0): H, W = arr.shape # 自定义分块大小,根据内存调整 block_h = 3000 block_w = 3000 out_h = (H + window_size - 1) // window_size out_w = (W + window_size - 1) // window_size result = np.zeros((out_h, out_w), dtype=arr.dtype) for i in range(0, H, block_h): for j in range(0, W, block_w): # 取包含当前块及相邻窗口的区域,避免边界窗口缺失数据 block = arr[i:i+block_h+window_size-1, j:j+block_w+window_size-1] # 计算当前块对应的输出区域坐标 out_i_start = i // window_size out_i_end = (i + block_h) // window_size out_j_start = j // window_size out_j_end = (j + block_w) // window_size # 对当前块计算多数投票 block_result = majority_vote_3x3(block, pad_val) # 写入结果数组 result[out_i_start:out_i_end, out_j_start:out_j_end] = block_result[:out_i_end-out_i_start, :out_j_end-out_j_start] return result
3. 结合Dask实现延迟计算
如果要继续用Xarray/Dask,将自定义函数包装为Dask可处理的形式,确保延迟执行:
import dask.array as da def dask_majority_vote(arr, window_size=3, pad_val=0): # 生成降采样后的分块规格 new_chunks = ((c//window_size for c in arr.chunks[0]), (c//window_size for c in arr.chunks[1])) return da.map_blocks( lambda block: majority_vote_3x3(block, pad_val), arr, chunks=new_chunks, dtype=arr.dtype ) # 使用示例:替换原数据集的变量 # ds['tree_type'] = dask_majority_vote(ds['tree_type'].data, window_size=3)
方案优势
np.bincount针对uint8类型高度优化,比scipy.stats.mode效率提升数倍- 数组重塑为视图操作,无额外内存拷贝
- 分块处理严格控制内存峰值,避免系统崩溃
内容的提问来源于stack exchange,提问作者dtm34
相关产品推荐
相关产品推荐

