np.where处理数组零元素除法触发警告?PyTorch无此问题
NumPy与PyTorch中where操作除零警告差异原因分析
核心差异在于参数计算逻辑
- NumPy的
np.where是预计算所有表达式:它会先完整执行b / a这个数组运算,不管a != 0的条件。当数组中存在a=0的元素时,除法运算会触发除零操作,进而抛出RuntimeWarning。只不过最后np.where会把这些错误计算的结果替换成0,所以最终输出结果是正确的。 - PyTorch的
torch.where是分支惰性计算:它会根据条件判断,只对满足a != 0的位置执行b / a运算,不满足条件的位置直接填充0,完全不会触碰a=0的元素进行除法,因此不会触发除零警告。
NumPy消除警告的可行方案
如果要避免NumPy的这个警告,有两种常用方式:
- 临时屏蔽除零警告(注意:会掩盖所有除零相关警告,谨慎使用):
import numpy as np np.seterr(divide='ignore') # 屏蔽除零警告
- 基于掩码的安全计算(更推荐,只计算有效位置):
a = np.array([[[4, 0,], [4, 4]], [[3, 3,], [3, 3]]]) b = np.ones((2, 2, 2)) result = np.zeros_like(b) mask = a != 0 result[mask] = b[mask] / a[mask] print(result)
测试代码与运行结果
测试代码
import torch import numpy as np # PyTorch 版本 a = torch.tensor([[[4, 0,], [4, 4]], [[3, 3,], [3, 3]]]) b = torch.ones((2, 2, 2)) b = torch.where(a != 0, b / a, 0) print(b) # NumPy 版本 a = np.array([[[4, 0,], [4, 4]], [[3, 3,], [3, 3]]]) b = np.ones((2, 2, 2)) b = np.where(a != 0, b / a, 0) print(b)
运行结果
tensor([[[0.2500, 0.0000], [0.2500, 0.2500]], [[0.3333, 0.3333], [0.3333, 0.3333]]], dtype=torch.float64) private/test.py:77: RuntimeWarning: divide by zero encountered in divide b = np.where(a != 0, b / a, 0) [[[0.25 0. ] [0.25 0.25 ]] [[0.33333333 0.33333333] [0.33333333 0.33333333]]]
内容的提问来源于stack exchange,提问作者Liu Tao
相关产品推荐
相关产品推荐

