如何用不同长度索引切片多维numpy数组并实现向量化元素更新
问题说明
现有如下数组:
import numpy as np R = np.array([[ 5., 3., 2., 7., 3., 6., 8., 9., 10., 55.], [ 5., 4., 2., 7., 3., 6., 8., 10., 10., 55.]]) F = np.array([[ 0.2 , 0.4 , 0.1 , 0.3 , 0.25, 0.25, 0.2 , 0.1 , 0.1 , 0.1 ], [ 0.3 , -0.4 , 0.1 , 0.3 , 0.25, 0.25, 0.4 , -0.4 , 0.1 , 0.1 ]]) K = np.array([[2], [1]])
需求:
- 对F的每一行分别排序,取排序后第一行前
K[0]个索引、第二行前K[1]个索引 - 用上述索引给R数组对应位置的元素加1,要求全程使用向量化操作,不使用for循环
你已经通过代码拿到了目标列索引:
indices = np.argsort(F)[np.tile(np.arange(F.shape[1]),(F.shape[0],1)) < K] # indices = np.array([7, 8, 7], dtype=int64)
解决方案
你已经拿到了列索引数组,只需要再构造匹配的行索引数组,用numpy原生的np.add.at做原地更新即可,全程无显式循环:
# 构造行索引:第一行需要更新K[0]个元素,对应行号均为0;第二行需要更新K[1]个元素,对应行号均为1 row_indices = np.repeat(np.arange(len(K)), K.flatten()) # row_indices 输出为 array([0, 0, 1]) # 复制原数组避免修改原始R Rnew = R.copy() # 向量化更新对应位置的元素,统一加1 np.add.at(Rnew, (row_indices, indices), 1)
验证输出:
print(Rnew) # 输出结果和预期完全一致: # [[ 5. 3. 2. 7. 3. 6. 8. 10. 11. 55.] # [ 5. 4. 2. 7. 3. 6. 8. 11. 10. 55.]]
np.add.at是numpy提供的无缓冲原地加法函数,即使同一个索引被多次选中也会正确累加,完全满足向量化操作要求。
内容的提问来源于stack exchange,提问作者PJORR
相关产品推荐
相关产品推荐

