Gets the device of the specified variable x if it is a tensor, or falls back to a default CPU device otherwise. Allows overriding by providing an explicit device. Args: x: a torch.Tensor to get the device from or another type device: Device (as str or torch.device)
(x, device: Optional[Device] = None)
| 35 | |
| 36 | |
| 37 | def get_device(x, device: Optional[Device] = None) -> torch.device: |
| 38 | """ |
| 39 | Gets the device of the specified variable x if it is a tensor, or |
| 40 | falls back to a default CPU device otherwise. Allows overriding by |
| 41 | providing an explicit device. |
| 42 | |
| 43 | Args: |
| 44 | x: a torch.Tensor to get the device from or another type |
| 45 | device: Device (as str or torch.device) to fall back to |
| 46 | |
| 47 | Returns: |
| 48 | A matching torch.device object |
| 49 | """ |
| 50 | |
| 51 | # User overrides device |
| 52 | if device is not None: |
| 53 | return make_device(device) |
| 54 | |
| 55 | # Set device based on input tensor |
| 56 | if torch.is_tensor(x): |
| 57 | return x.device |
| 58 | |
| 59 | # Default device is cpu |
| 60 | return torch.device("cpu") |
no test coverage detected