Numpy如何不使用花式索引按指定轴和索引设置ndarray元素值
numpy 原生没有完全匹配你所需用法的set_value内置方法,但你可以通过构造广播兼容的索引元组,自行封装实现该功能,代码如下:
实现方案
首先封装独立赋值函数:
import numpy as np def set_value(arr, value, axis, indices): axis = tuple(axis) indices = tuple(indices) # 初始化所有轴为全选状态 idx = [slice(None)] * arr.ndim for ax, idx_list in zip(axis, indices): # 将当前轴的索引扩展维度,适配多轴广播规则 expanded_idx = np.array(idx_list) # 左侧补维度 for _ in range(ax): expanded_idx = np.expand_dims(expanded_idx, axis=0) # 右侧补维度 for _ in range(ax + 1, arr.ndim): expanded_idx = np.expand_dims(expanded_idx, axis=-1) idx[ax] = expanded_idx arr[tuple(idx)] = value
使用示例
完全匹配你预期的调用逻辑:
z = np.zeros((3,4,2,8)) # 为轴0的[1,2]位置、轴2的[0,1]位置赋值为9 set_value(z, 9, axis=(0,2), indices=[[1,2],[0,1]]) # 验证和原花式索引效果一致 print(np.all(z[[1,2],:,[0,1],:] == 9)) # 输出 True
如果想要实现z.set_value的调用形式,可以给ndarray动态绑定方法:
np.ndarray.set_value = lambda self, value, axis, indices: set_value(self, value, axis, indices) # 调用效果 z = np.zeros((3,4,2,8)) z.set_value(9, axis=(0,2), indices=[[1,2],[0,1]])
补充说明
你提到的np.put是按数组扁平化后的一维索引赋值,np.put_along_axis仅支持单轴赋值,确实无法直接满足多轴同时指定索引赋值的需求,上面的实现中用expand_dims做维度扩展就是为了让多轴的索引可以互相广播适配数组形状,实现和花式索引完全一致的效果。
内容的提问来源于stack exchange,提问作者stevemo
相关产品推荐
相关产品推荐

