如何在Dask中实现类NumPy的广播式索引切片?
在Dask数组中实现按指定索引数组提取元素的方法
在NumPy里,我们可以直接用二维索引数组从目标数组中提取对应位置的元素,比如下面这段代码能正常运行:
import numpy as np x = np.linspace(0,5,10).reshape(10,1) # 形状为(1,3)的索引数组 filt = np.array([[2,3,5]]) # 提取索引为2、3、5的元素 x[filt]
但换成Dask数组时,直接用x[filt]会触发AssertionError,这里提供两种可行的解决方式:
方法一:使用da.take
da.take是Dask专门用来按索引提取元素的方法,先把索引数组展平成一维,提取后再调整回原索引数组的形状:
import dask.array as da x = da.linspace(0,5,10).reshape(10,1) filt = da.array([[2,3,5]]) # 提取元素并重塑形状 result = da.take(x, filt.flatten()).reshape(filt.shape) # 计算并输出结果 print(result.compute())
方法二:展平索引后再重塑
先将二维的索引数组展平为一维,用它提取元素后,再把结果重塑成和原索引数组一致的形状:
import dask.array as da x = da.linspace(0,5,10).reshape(10,1) filt = da.array([[2,3,5]]) # 展平索引提取元素,再调整形状 result = x[filt.flatten()].reshape(filt.shape) print(result.compute())
为什么原方法会报错?
Dask数组的索引规则和NumPy有差异,NumPy支持用二维数组对一维数组做广播式索引,但Dask不支持这种操作,所以需要先将索引数组转为一维,提取完成后再恢复目标形状。
内容的提问来源于stack exchange,提问作者matsuo_basho
相关产品推荐
相关产品推荐

