Spec-Zone.ru › PyTorch 2.14

torch.cuda.jiterator._create_jit_fn

torch.cuda.jiterator._create_jit_fn(code_string, **kwargs) [исходный код]

Создает CUDA-ядро для поэлементной операции, сгенерированное jiterator.

Строка кода должна содержать корректную функцию CUDA, описывающую вычисление для одного элемента. Строка кода должна соответствовать шаблону C++, как показано в примере ниже. Эта функция будет встроена в шаблон поэлементного ядра и скомпилирована на лету. Скомпилированное ядро будет кэшироваться в памяти, а также во временном локальном каталоге.

Ядра, сгенерированные jiterator, поддерживают несмежные тензоры, широковещание и повышение типов.

Параметры:
  • code_string (str) – Строка кода CUDA для компиляции с помощью jiterator. Функтор точки входа должен возвращать значение.
  • kwargs (Dict, optional) – Именованные аргументы для сгенерированной функции
Тип возвращаемого значения:

Callable

Пример:

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.0

Jiterator можно использовать совместно с регистрацией в 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.

© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.cuda.jiterator._create_jit_fn.html

Spec-Zone.ru

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