如何在3D数组中沿单轴计算连续1的序列长度(游程编码相关)
3D数组沿单轴计算连续1的序列长度(无循环实现)
需求背景
我需要在3D(纬度×经度×时间)的xarray数据中,沿时间轴计算每个网格单元上连续1的序列长度。输出结果需和原数组维度一致:仅在连续1序列的起始位置标注序列长度,其余位置填充NaN。比如1D数组[0,1,0,0,1,1,1,0,1,1]对应的结果是[nan,1,nan,nan,3,nan,nan,nan,2,nan]。
由于气候科学研究中经纬度网格数量极大,循环遍历每个网格单元效率极低,因此需要无循环的向量化实现。目前已实现1D版本,但扩展到3D时卡在了有效元素的筛选与前置处理上,求可行的无循环解决方案。
已实现的1D解决方案
import xarray as xr import numpy as np def run_lengths(da): n = len(da) # 标记值发生变化的位置 y = da.values[1:] != da.values[:-1] # 获取所有变化点的索引,加上最后一个元素的索引 i = np.append(np.where(y), n - 1) # 计算每个连续段的长度 z = np.diff(np.append(-1, i)) # 每个连续段的起始索引 p = np.cumsum(np.append(0, z))[:-1] # 筛选出连续1的段 runs = np.where(da[i] == 1)[0] runs_len = z[runs] # 连续1的序列长度 time_val = da.time[p[runs]] # 序列起始时间 # 创建结果数组并对齐时间轴 da_runs = xr.DataArray(runs_len, coords={'time': time_val}) _, da_runs = xr.align(da, da_runs, join='outer') return da_runs # 测试1D情况 da = xr.DataArray( np.array([[[0,1,1,0,0,0],[1,0,1,1,0,1],[1,1,1,1,0,1]],[[0,1,1,0,0,0],[1,0,1,1,0,1],[1,1,1,1,0,1]]]), coords={'lat': [0,1], 'lon': [0,1,2], 'time': [0,1,2,3,4,5]} ) da_runs = run_lengths(da[0,1]) print(da_runs)
输出结果:
<xarray.DataArray (time: 6)> array([ 1., nan, 2., nan, nan, 1.]) Coordinates: * time (time) int64 0 1 2 3 4 5
3D无循环实现方案
以下是基于向量化操作的3D版本实现,核心思路是在整个3D数组上批量处理所有网格的连续段,避免循环:
import xarray as xr import numpy as np def run_lengths_3D(da): # 获取维度信息 lat_dim, lon_dim, time_dim = da.dims n_time = da.sizes[time_dim] lat_vals = da.coords[lat_dim].values lon_vals = da.coords[lon_dim].values time_vals = da.coords[time_dim].values # 将数据转为numpy数组(lat, lon, time) data = da.values # 1. 标记每个网格点上时间序列的变化位置(time维度上的差分) # 变化点标记为True,shape: (lat, lon, time-1) change_mask = data[..., 1:] != data[..., :-1] # 在每个序列末尾添加一个True(确保最后一段被捕获) change_mask = np.concatenate([change_mask, np.ones((data.shape[0], data.shape[1], 1), dtype=bool)], axis=-1) # 2. 获取每个网格点上所有连续段的起始和结束索引 # 生成时间轴索引数组,shape: (1,1,time) time_indices = np.arange(n_time).reshape(1, 1, -1) # 每个网格点上变化点的索引,用NaN填充非变化点 change_indices = np.where(change_mask, time_indices, np.nan) # 3. 计算每个连续段的长度和起始位置 # 沿时间轴向前填充NaN,得到每个位置所属连续段的结束索引 segment_end = np.full_like(change_indices, np.nan) for t in range(n_time-1, -1, -1): segment_end[..., t] = np.where(~np.isnan(change_indices[..., t]), change_indices[..., t], segment_end[..., t+1] if t+1 < n_time else np.nan) # 计算每个连续段的起始索引:前一个段的结束索引+1,首段起始为0 segment_start = np.roll(segment_end, shift=1, axis=-1) segment_start[..., 0] = 0 # 修正起始索引:非变化点的起始索引等于前一个位置的起始索引 for t in range(1, n_time): segment_start[..., t] = np.where(change_mask[..., t-1], segment_start[..., t], segment_start[..., t-1]) # 连续段长度 segment_length = segment_end - segment_start + 1 # 4. 筛选出连续1的段,其余位置设为NaN # 检查每个段的第一个值是否为1 is_one_segment = data[..., 0] == 1 # 首段的第一个值 # 对于非首段,检查段起始位置的值 for t in range(1, n_time): is_one_segment = np.where(change_mask[..., t-1], data[..., t] == 1, is_one_segment) # 最终结果:仅在段起始位置保留长度,其余为NaN result = np.where((time_indices == segment_start) & is_one_segment, segment_length, np.nan) # 转为xarray DataArray da_result = xr.DataArray( result, coords={lat_dim: lat_vals, lon_dim: lon_vals, time_dim: time_vals}, dims=(lat_dim, lon_dim, time_dim) ) return da_result # 测试3D情况 da = xr.DataArray( np.array([[[0,1,1,0,0,0],[1,0,1,1,0,1],[1,1,1,1,0,1]],[[0,1,1,0,0,0],[1,0,1,1,0,1],[1,1,1,1,0,1]]]), coords={'lat': [0,1], 'lon': [0,1,2], 'time': [0,1,2,3,4,5]} ) da_runs_3D = run_lengths_3D(da) print(da_runs_3D)
方案说明
- 变化点标记:通过时间轴上的差分操作,批量标记所有网格点上数值发生变化的位置,确保每个连续段的结束位置被捕获。
- 连续段信息计算:通过向后填充变化点索引,得到每个位置所属连续段的结束索引,再推导起始索引和段长度。
- 筛选连续1段:标记出所有由1组成的连续段,仅在段的起始位置保留长度值,其余位置设为
NaN。 - 向量化操作:所有计算均基于numpy的数组广播机制,避免了对经纬度网格的循环遍历,大幅提升大网格数据的处理效率。
替代方案:使用xarray的groupby与cumsum
另一种更简洁的实现方式,利用xarray的分组功能:
def run_lengths_3D_groupby(da): # 标记每个网格点上连续段的分组ID # 当值变化时,分组ID+1 change_mask = da.diff(dim='time') != 0 group_id = change_mask.cumsum(dim='time').fillna(0) # 首时间步的分组ID设为0 group_id = xr.concat([xr.DataArray(np.zeros_like(da.isel(time=0)), coords=da.isel(time=0).coords), group_id], dim='time') # 计算每个分组的长度,以及分组的起始时间 group_length = da.groupby(group_id).count(dim='time') group_first_time = da.groupby(group_id).first(dim='time') # 筛选出值为1的分组 valid_groups = group_first_time.where(group_first_time == 1, drop=True) valid_lengths = group_length.sel(group_id=valid_groups.group_id) # 创建结果数组,仅在分组起始时间填充长度 result = xr.full_like(da, np.nan) result.loc[{'time': valid_groups.time}] = valid_lengths.values return result
这个方案利用xarray的groupby自动处理多维分组,代码更简洁,同样无需循环,适合xarray用户的使用习惯。
内容的提问来源于stack exchange,提问作者KateW12
相关产品推荐
相关产品推荐

