torch.cuda.jiterator._create_jit_fn
-
torch.cuda.jiterator._create_jit_fn(code_string, **kwargs)[source] -
Создать ядро CUDA, сгенерированное jiterator, для оператора по элементам.
Строка кода должна быть допустимой функцией CUDA, описывающей вычисление для одного элемента. Строка кода должна следовать шаблону шаблона C++, как показано в примере ниже. Эта функция будет встроена в шаблон ядра по элементам и скомпилирована на лету. Скомпилированное ядро будет кэшироваться в памяти, а также в локальном временном каталоге.
Ядра, сгенерированные jiterator, принимают несмежные тензоры и поддерживают распределение и повышение типа.
- Параметры:
-
- code_string (str) – Строка кода CUDA, которая будет скомпилирована jiterator. Входной функтор должен возвращать значение по значению.
- kwargs (Dict, optional) – Параметры ключевого слова для сгенерированной функции
- Тип возвращаемого значения:
Пример:
code_string = "template <typename T> T my_kernel(T x, T y, T alpha) { return -x + alpha * y; }" jitted_fn = create_jit_fn(code_string, alpha=1.0) a = torch.rand(3, device='cuda') b = torch.rand(3, device='cuda') # invoke jitted function like a regular python function result = jitted_fn(a, b, alpha=3.14)code_string также позволяет определять несколько функций, и последняя функция будет рассматриваться как входная функция.
Пример:
code_string = "template <typename T> T util_fn(T x, T y) { return ::sin(x) + ::cos(y); }" code_string += "template <typename T> T my_kernel(T x, T y, T val) { return ::min(val, util_fn(x, y)); }" jitted_fn = create_jit_fn(code_string, val=0.0) a = torch.rand(3, device='cuda') b = torch.rand(3, device='cuda') # invoke jitted function like a regular python function result = jitted_fn(a, b) # using default val=0.0Jiterator может использоваться совместно с регистрацией Python для переопределения ядра CUDA оператора. Следующий пример переопределяет ядро CUDA gelu с помощью relu.
Пример:
code_string = "template <typename T> T my_gelu(T a) { return a > 0 ? a : 0; }" my_gelu = create_jit_fn(code_string) my_lib = torch.library.Library("aten", "IMPL") my_lib.impl('aten::gelu', my_gelu, "CUDA") # torch.nn.GELU and torch.nn.function.gelu are now overridden a = torch.rand(3, device='cuda') torch.allclose(torch.nn.functional.gelu(a), torch.nn.functional.relu(a))Предупреждение
Этот API находится в стадии бета-тестирования и может быть изменён в будущих выпусках.
Предупреждение
Этот API поддерживает только до 8 входных данных и 1 выходного.
Предупреждение
Все входные тензоры должны находиться в устройстве CUDA.
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.cuda.jiterator._create_jit_fn.html