如何使用Dask重构从三维数组维度子集获取的子数组?
Dask布尔索引赋值重构数组报错解决方法
问题背景
现有三维NumPy数组,基于前两个维度的布尔条件拆分得到形状为(5,4)的arr1和(1,4)的arr2。尝试用Dask创建与原数组形状一致的零数组并通过布尔索引赋值重构时,触发错误:
ValueError: Boolean index assignment in Dask expects equally shaped arrays
报错原因
Dask的布尔索引赋值要求索引数组与被赋值数组的形状完全匹配。原代码中bool_check是(2,3)的二维数组,而目标数组arr3是(2,3,4)的三维数组,两者维度不匹配,无法直接对应赋值。
解决方案
方法1:使用da.where实现条件赋值
通过扩展布尔条件到三维,结合da.where选择对应的值完成重构:
import dask.array as da import numpy as np np.random.seed(40) test_arr = np.random.normal(size=(2,3,4)) bool_check = test_arr[:,:,0] < 0.6 # 转换为Dask数组 da_bool = da.from_array(bool_check) da_arr1 = da.from_array(test_arr[bool_check]) da_arr2 = da.from_array(test_arr[~bool_check]) # 将二维布尔条件扩展为三维(匹配原数组形状) expanded_bool = da.broadcast_to(da_bool[..., None], test_arr.shape) # 构造与原形状匹配的填充数组 def fill_to_shape(bool_mask, src_arr, target_shape): # 先填充前两维 temp = da.zeros(target_shape[:2], dtype=src_arr.dtype) idx = da.where(bool_mask) temp[idx] = src_arr.reshape(-1) # 扩展到第三维 return da.broadcast_to(temp[..., None], target_shape) filled_arr1 = fill_to_shape(da_bool, da_arr1, test_arr.shape) filled_arr2 = fill_to_shape(~da_bool, da_arr2, test_arr.shape) # 用where完成重构 arr3 = da.where(expanded_bool, filled_arr1, filled_arr2) # 验证结果 print(arr3.compute())
方法2:通过整数索引赋值
将布尔条件转换为整数索引,展平后对应赋值:
import dask.array as da import numpy as np np.random.seed(40) test_arr = np.random.normal(size=(2,3,4)) bool_check = test_arr[:,:,0] < 0.6 da_bool = da.from_array(bool_check) da_arr1 = da.from_array(test_arr[bool_check]) da_arr2 = da.from_array(test_arr[~bool_check]) arr3 = da.zeros_like(test_arr) # 处理arr1的赋值:获取布尔条件对应的整数索引并扩展到第三维 idx = da.where(da_bool) # 构造三维索引:每个前两维索引对应第三维的所有位置 idx_3d = (idx[0][:, None], idx[1][:, None], da.arange(test_arr.shape[2])) # 展平索引和待赋值数组,完成赋值 arr3[(idx_3d[0].flatten(), idx_3d[1].flatten(), idx_3d[2].flatten())] = da_arr1.flatten() # 处理arr2的赋值 idx_not = da.where(~da_bool) idx_not_3d = (idx_not[0][:, None], idx_not[1][:, None], da.arange(test_arr.shape[2])) arr3[(idx_not_3d[0].flatten(), idx_not_3d[1].flatten(), idx_not_3d[2].flatten())] = da_arr2.flatten() # 验证结果 print(arr3.compute())
内容的提问来源于stack exchange,提问作者matsuo_basho
相关产品推荐
相关产品推荐

