Files
pytorch/torch/accelerator/_utils.py
Dmitry Rogozhkin 7314cf44ae torch/accelerator: fix device type comparison (#143541)
This was failing without the fix:
```
python -c 'import torch; d=torch.device("xpu:0"); torch.accelerator.current_stream(d)'
```
with:
```
ValueError: xpu doesn't match the current accelerator xpu.
```

CC: @guangyey, @EikanWang

Pull Request resolved: https://github.com/pytorch/pytorch/pull/143541
Approved by: https://github.com/guangyey, https://github.com/albanD
2024-12-23 10:54:53 +00:00

26 lines
907 B
Python

from typing import Optional
import torch
from torch.types import Device as _device_t
def _get_device_index(device: _device_t, optional: bool = False) -> int:
if isinstance(device, int):
return device
if isinstance(device, str):
device = torch.device(device)
device_index: Optional[int] = None
if isinstance(device, torch.device):
if torch.accelerator.current_accelerator().type != device.type:
raise ValueError(
f"{device.type} doesn't match the current accelerator {torch.accelerator.current_accelerator()}."
)
device_index = device.index
if device_index is None:
if not optional:
raise ValueError(
f"Expected a torch.device with a specified index or an integer, but got:{device}"
)
return torch.accelerator.current_device_index()
return device_index