Fill invalid values (NaN and infinite) in a NumPy array or PyTorch tensor with a specified fill value. Args: arr: fill_value: Returns: Union[np.ndarray, torch.Tensor]: Processed array or tensor with invalid values replaced by fill_value.
(arr: Union[np.ndarray, torch.Tensor], fill_value: float = 0.0)
| 23 | return arr |
| 24 | |
| 25 | def fill_invalid_values(arr: Union[np.ndarray, torch.Tensor], fill_value: float = 0.0) -> Union[np.ndarray, torch.Tensor]: |
| 26 | """ |
| 27 | Fill invalid values (NaN and infinite) in a NumPy array or PyTorch tensor with a specified fill value. |
| 28 | Args: |
| 29 | arr: |
| 30 | fill_value: |
| 31 | |
| 32 | Returns: |
| 33 | Union[np.ndarray, torch.Tensor]: Processed array or tensor with invalid values replaced by fill_value. |
| 34 | """ |
| 35 | |
| 36 | if isinstance(arr, torch.Tensor): |
| 37 | arr = torch.nan_to_num(arr, nan=.0, posinf=1.0, neginf=.0) |
| 38 | elif isinstance(arr, np.ndarray): |
| 39 | arr = np.nan_to_num(arr, nan=.0, posinf=1.0, neginf=.0) |
| 40 | else: |
| 41 | raise TypeError("Input must be a NumPy array or a PyTorch tensor.") |
| 42 | return arr |