Recover input type to original array type. Args: input_array (np.ndarray or torch.Tensor): Input array. Returns: np.ndarray or torch.Tensor or int or float: Converted array.
(
self, input_array: Union[np.ndarray, torch.Tensor]
)
| 324 | return converted_array |
| 325 | |
| 326 | def recover( |
| 327 | self, input_array: Union[np.ndarray, torch.Tensor] |
| 328 | ) -> Union[np.ndarray, torch.Tensor, int, float]: |
| 329 | """Recover input type to original array type. |
| 330 | |
| 331 | Args: |
| 332 | input_array (np.ndarray or torch.Tensor): Input array. |
| 333 | |
| 334 | Returns: |
| 335 | np.ndarray or torch.Tensor or int or float: Converted array. |
| 336 | """ |
| 337 | assert isinstance(input_array, (np.ndarray, torch.Tensor)), \ |
| 338 | 'invalid input array type' |
| 339 | if isinstance(input_array, self.array_type): |
| 340 | return input_array |
| 341 | elif isinstance(input_array, torch.Tensor): |
| 342 | converted_array = input_array.cpu().numpy().astype(self.dtype) |
| 343 | else: |
| 344 | converted_array = torch.tensor(input_array, |
| 345 | dtype=self.dtype, |
| 346 | device=self.device) |
| 347 | if self.is_num: |
| 348 | converted_array = converted_array.item() |
| 349 | return converted_array |
no test coverage detected