PyTorch中除法遇零产生NaN时,如何将结果替换为0?
Got it, let's solve this problem. When you divide tensors where the denominator has zeros in PyTorch, you end up with NaNs in those positions—and you want to replace those NaNs with 0s. Here are three straightforward, efficient ways to do this, using your sample code as a starting point:
First, let's recap your original code (with imports added for completeness):
import torch as th import numpy as np a = th.from_numpy(np.array([ [1, 0], [0, 1], [1, 1]])) b = th.zeros_like(a) b[0, :] = 2 # This will produce NaNs where b is 0 original_result = a / b print(original_result) # Output: # tensor([[0.5000, 0.0000], # [ nan, nan], # [ nan, nan]], dtype=torch.float64)
1. Use torch.nan_to_num (Recommended)
This is the cleanest and most efficient method, especially if you're using PyTorch 1.8 or newer. The nan_to_num function is purpose-built to replace NaNs (and infinities) with specified values.
fixed_result = th.nan_to_num(a / b, nan=0.0) print(fixed_result) # Output: # tensor([[0.5000, 0.0000], # [0.0000, 0.0000], # [0.0000, 0.0000]], dtype=torch.float64)
Just pass nan=0.0 to tell PyTorch to replace all NaN values with 0. You can also optionally handle positive/negative infinities with the posinf and neginf parameters if needed.
2. Pre-empt with torch.where
If you want to avoid generating NaNs entirely, use torch.where to check for zero denominators before performing the division. This conditionally returns 0 where b is 0, and the division result elsewhere.
fixed_result = th.where(b == 0, th.tensor(0.0, device=a.device), a / b) print(fixed_result)
Make sure the tensor you use for the "0 case" matches the device (CPU/GPU) and data type of your original tensors to avoid errors.
3. Manually Replace NaNs Post-Division
If you prefer a more explicit approach, first compute the division, then locate and replace the NaNs directly:
fixed_result = a / b # Mask out NaN positions and set them to 0 fixed_result[fixed_result.isnan()] = 0 print(fixed_result)
This works well for simple cases, though for large tensors, it's slightly less efficient than the first two methods since you're creating an extra mask tensor.
All three methods will give you the same desired output—replacing those pesky NaNs with 0s. Pick the one that fits your code style and PyTorch version best!
内容的提问来源于stack exchange,提问作者GoingMyWay

