torch.library
API регистрации операторов Python предоставляет возможности расширения основного библиотеки операторов PyTorch с пользовательскими операторами. В настоящее время это можно сделать двумя способами:
-
Создание новых библиотек
-
Позволяет регистрировать новые операторы и ядра для различных бэкэндов и функциональных возможностей, указав соответствующие ключи диспетчеризации. Например,
- Рассмотрим регистрацию нового оператора
addв вашем недавно созданном пространстве именfoo. Вы можете получить доступ к этому оператору, используя APItorch.opsи вызывая его, вызываяtorch.ops.foo.add. Вы также можете получить доступ к конкретным зарегистрированным перегрузкам, вызвавtorch.ops.foo.add.{overload_name}. - Если вы зарегистрировали новое ядро для ключа диспетчеризации
CUDAдля этого оператора, тогда ваша настраиваемая функция будет вызвана для входных тензоров CUDA.
- Рассмотрим регистрацию нового оператора
- Это можно сделать, создав объекты класса Library типа
"DEF".
-
-
Расширение существующих C++ библиотек (например, aten)
- Позволяет регистрировать ядра для существующих операторов, соответствующих различным бэкэндам и функциональным возможностям, указав соответствующие ключи диспетчеризации.
-
Это может быть полезно для заполнения пробелов в поддержке операторов для функции, реализованной с помощью ключа диспетчеризации. Например,
- Вы можете добавить поддержку оператора для мета-тензоров (зарегистрировав функцию для ключа диспетчеризации
Meta).
- Вы можете добавить поддержку оператора для мета-тензоров (зарегистрировав функцию для ключа диспетчеризации
- Это можно сделать, создав объекты класса Library типа
"IMPL".
Руководство, которое проведет вас через некоторые примеры использования этого API, доступно на Google Colab.
Предупреждение
Диспетчер — это сложная концепция PyTorch, и глубокое понимание Диспетчера имеет решающее значение для выполнения любых сложных задач с этим API. Эта публикация блога является хорошей отправной точкой для изучения Диспетчера.
-
class torch.library.Library(ns, kind, dispatch_key='')[source] -
Класс для создания библиотек, которые могут использоваться для регистрации новых операторов или переопределения операторов в существующих библиотеках из Python. Пользователь может необязательно передать имя ключа диспетчеризации, если он хочет регистрировать ядра, соответствующие только одному конкретному ключу диспетчеризации.
Чтобы создать библиотеку для переопределения операторов в существующей библиотеке (с именем ns), установите тип в «IMPL». Чтобы создать новую библиотеку (с именем ns) для регистрации новых операторов, установите тип в «DEF». Чтобы создать фрагмент, возможно, существующей библиотеки для регистрации операторов (и обойти ограничение, что для данного пространства имен существует только одна библиотека), установите тип в «FRAGMENT».
- Параметры
-
- ns – имя библиотеки
- kind – «DEF», «IMPL» (по умолчанию: «IMPL»), «FRAGMENT»
- dispatch_key – ключ диспетчеризации PyTorch (по умолчанию: «»)
-
define(schema, alias_analysis='')[source] -
Определяет новый оператор и его семантику в пространстве имен ns.
- Параметры
-
- schema – схема функции для определения нового оператора.
- alias_analysis (необязательно) – Указывает, могут ли быть выведены свойства алиасов аргументов оператора из схемы (по умолчанию) или нет («CONSERVATIVE»).
- Возвращает
-
имя оператора, как выведено из схемы.
- Пример::
-
>>> my_lib = Library("foo", "DEF") >>> my_lib.define("sum(Tensor self) -> Tensor")
-
impl(op_name, fn, dispatch_key='')[source] -
Регистрирует реализацию функции для оператора, определенного в библиотеке.
- Параметры
-
- op_name – имя оператора (вместе с перегрузкой) или объект OpOverload.
-
fn – функция, которая является реализацией оператора для входного ключа диспетчеризации или
fallthrough_kernel()для регистрации пропуска. - dispatch_key – ключ диспетчеризации, для которого должна быть зарегистрирована входная функция. По умолчанию используется ключ диспетчеризации, с которым была создана библиотека.
- Пример::
-
>>> my_lib = Library("aten", "IMPL") >>> def div_cpu(self, other): >>> return self * (1 / other) >>> my_lib.impl("div.Tensor", div_cpu, "CPU")
-
torch.library.fallthrough_kernel()[source] -
Функция-заглушка для передачи в
Library.implдля регистрации пропуска.
Мы также добавили несколько декораторов функций, чтобы упростить регистрацию функций для операторов:
torch.library.impl()torch.library.define()
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/library.html