Spec-Zone.ru › PyTorch 2

torch.utils.rename_privateuse1_backend

torch.utils.rename_privateuse1_backend(backend_name) → None [source]

Этот API следует использовать для переименования устройства частного использования 1, чтобы его удобнее было использовать в качестве имени устройства в API PyTorch.

Шаги:

  1. (В C++) реализуйте ядра для различных операций torch и зарегистрируйте их в ключе распределения PrivateUse1.
  2. (В Python) вызовите torch.utils.rename_privateuse1_backend(“foo”)

Теперь вы можете использовать «foo» как обычную строку устройства в Python.

Примечание: этот API можно вызывать только один раз на процесс. Попытка изменить внешнее устройство после его установки приведет к ошибке.

Примечание (AMP): Если вы хотите поддерживать AMP на своём устройстве, вы можете зарегистрировать модуль пользовательского бэкэнда. Бэкэнд должен зарегистрировать модуль пользовательского бэкэнда с torch._register_device_module("foo", BackendModule). Модуль Backend должен иметь следующие API:

  1. get_amp_supported_dtype() -> List[torch.dtype] получить поддерживаемые типы данных на вашем устройстве «foo» в AMP, возможно, устройство «foo» поддерживает ещё один тип данных.
  2. is_autocast_enabled() -> bool проверить, включен ли AMP на вашем устройстве «foo».
  3. get_autocast_dtype() -> torch.dtype получить поддерживаемый тип данных на вашем устройстве «foo» в AMP, который устанавливается set_autocast_dtype, или по умолчанию, а тип данных по умолчанию — torch.float16.
  4. set_autocast_enabled(bool) -> None включить или отключить AMP на вашем устройстве «foo».
  5. set_autocast_dtype(dtype) -> None установить поддерживаемый тип данных на вашем устройстве «foo» в AMP, и тип данных должен содержаться в типах данных, полученных из get_amp_supported_dtype.

Примечание (случайное): если вы хотите иметь возможность устанавливать seed для вашего устройства, модуль Backend должен иметь следующие API:

  1. _is_in_bad_fork() -> bool Возвращает True, если сейчас он находится в bad_fork, иначе возвращает False.
  2. manual_seed_all(seed int) -> None Устанавливает seed для генерации случайных чисел для ваших устройств.
  3. device_count() -> int Возвращает количество доступных «foo».
  4. get_rng_state(device: Union[int, str, torch.device] = 'foo') -> Tensor Возвращает список ByteTensor, представляющих состояния генераторов случайных чисел для всех устройств.
  5. set_rng_state(new_state: Tensor, device: Union[int, str, torch.device] = 'foo') -> None Устанавливает состояние генератора случайных чисел указанного устройства «foo».

И есть некоторые общие функции:

  1. is_available() -> bool Возвращает значение булевого типа, указывающее, доступно ли устройство «foo» в данный момент.
  2. current_device() -> int Возвращает индекс выбранного устройства.

Для получения более подробной информации см. https://pytorch.org/tutorials/advanced/extend_dispatcher.html#get-a-dispatch-key-for-your-backend Для примера существующего решения см. https://github.com/bdhirsh/pytorch_open_registration_example

Пример:

>>> torch.utils.rename_privateuse1_backend("foo")
# This will work, assuming that you've implemented the right C++ kernels
# to implement torch.ones.
>>> a = torch.ones(2, device="foo")

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.utils.rename_privateuse1_backend.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API