如何使用np.where根据条件修改二维ndarray中的值?
解决NumPy数组修改不生效的问题
看起来你踩了NumPy里数组视图与副本的坑,还有布尔运算优先级的小错误,咱们一步步捋清楚解决:
原代码的问题分析
你的代码里有两个关键问题导致原数组没变化:
- 布尔运算优先级错误:
&的优先级比<高,所以arr[1] < 0 & arr[2] <0会被解析成arr[1] < (0 & arr[2]),完全偏离了你想要的逻辑,必须给每个条件单独加括号:(arr[1] < 0) & (arr[2] < 0) - 操作的是数组副本而非视图:
arr[:, np.where(...)]返回的是原数组的副本,你修改这个副本的内容只会改变副本,原数组自然不会有变化。
正确的实现方式
我们可以直接用布尔索引操作原数组的视图,这样修改会直接作用于原数组:
import numpy as np # 先修正你示例数组的长度不一致问题(原第一行5个元素,后两行10个,这里统一为10个元素演示) arr = np.array([ [5., 6., 7., 8., 9., 10., 11., 12., 13., 14.], [-0.6810069, -0.61737489, 0.09869664, -0.95659638, 0.54052288, -0.21486195, -0.55328357, -0.41680664, 0.60628816, -0.92563772], [0.54264171, 0.52459383, -0.83691756, 0.66207278, -0.65591567, -0.72713526, -0.66719761, -0.99448398, 0.06691338, -0.2620483] ]) # 定义两个条件掩码 mask_add = (arr[1] < 0) & (arr[2] < 0) # 满足加10的条件 mask_sub = (arr[1] < 0) & (arr[2] > 0) # 满足减10的条件 # 直接修改原数组的第一行对应位置 arr[0, mask_add] += 10 arr[0, mask_sub] -= 10 print(arr[0])
运行后原数组会直接被修改,第一行的输出结果为:
[-5. -4. 7. -2. 9. 20. 21. 22. 13. 24.]
为什么这样可行?
布尔索引返回的是原数组的视图(不是副本),对视图的修改会直接同步到原数组。而np.where返回的是索引元组,用这些非连续索引切片时,NumPy会创建副本,导致修改无法同步到原数组。
如果你的实际数组确实存在行长度不一致的情况,需要先统一数组维度,否则会触发维度不匹配的报错哦。
内容的提问来源于stack exchange,提问作者Wrenb
相关产品推荐
相关产品推荐

