报错修正:用作索引的数组需为整数/布尔类型的NumPy代码问题
问题解决:按索引提取numpy数组元素
错误根源
你的np.apply_along_axis用法完全错误——这个函数传入lambda的x是luck_A沿axis=2切割后的浮点数元素序列,不是维度索引。拿浮点数数组去索引其他数组,必然触发类型错误,哪怕转int也没用,因为逻辑方向就错了。
正确解法(推荐)
用numpy原生的np.take_along_axis函数,专门处理这种“沿指定轴按索引提取元素”的场景,代码简洁且高效:
import numpy as np np.random.seed(0) luck_A = np.random.uniform(0, 1, size=(2, 1, 100)) A = np.random.uniform(0, 1, size=(2, 1, 100)) # 获取每个(2,1)维度中Top10元素的索引 index_select_A = np.argpartition(A, -10)[:, :, -10:] # 沿axis=2提取对应索引的元素 result = np.take_along_axis(luck_A, index_select_A, axis=2) # 验证结果形状:符合预期的(2,1,10) print(result.shape) # 输出 (2, 1, 10)
手动构造索引的替代方案
如果不想用take_along_axis,可以手动构造前两个维度的索引数组,通过广播匹配index_select_A的形状:
# 构造前两个维度的索引,和index_select_A广播对齐 i = np.arange(luck_A.shape[0])[:, None, None] # 形状(2,1,1) j = np.arange(luck_A.shape[1])[None, :, None] # 形状(1,1,1) result = luck_A[i, j, index_select_A]
内容的提问来源于stack exchange,提问作者llexiss
相关产品推荐
相关产品推荐

