Ограничения пользовательского интерфейса
Создано: 12 июня 2025 г. | Последнее обновление: 12 июня 2025 г.
torch.func, как и JAX, имеет ограничения на преобразуемые функции. В целом ограничения JAX заключаются в том, что преобразования работают только с чистыми функциями: то есть функциями, выходные данные которых полностью определяются входными данными и которые не имеют побочных эффектов (например, мутаций).
У нас действует аналогичное правило: наши преобразования хорошо работают с чистыми функциями. Однако мы поддерживаем некоторые операции на месте. С одной стороны, для работы с преобразованиями функций может потребоваться изменить способ написания кода PyTorch; с другой стороны, наши преобразования позволяют выражать вещи, которые раньше было сложно выразить в PyTorch.
Общие ограничения
Все преобразования torch.func имеют общее ограничение: функция не должна присваивать значения глобальным переменным. Вместо этого все выходные данные функции должны возвращаться из нее. Это ограничение связано с тем, как реализован torch.func: каждое преобразование оборачивает входные тензоры в специальные подклассы тензоров torch.func, которые обеспечивают работу преобразования.
Итак, вместо следующего:
import torch from torch.func import grad # Don't do this intermediate = None def f(x): global intermediate intermediate = x.sin() z = intermediate.sin() return z x = torch.randn([]) grad_x = grad(f)(x)
Перепишите f так, чтобы она возвращала intermediate:
def f(x): intermediate = x.sin() z = intermediate.sin() return z, intermediate grad_x, intermediate = grad(f, has_aux=True)(x)
API torch.autograd
Если вы пытаетесь использовать API torch.autograd, например torch.autograd.grad или torch.autograd.backward, внутри функции, преобразуемой с помощью vmap() или одного из AD-преобразований torch.func (vjp(), jvp(), jacrev(), jacfwd()), преобразование может оказаться не в состоянии применить преобразование к этому API. Если это не удастся, вы получите сообщение об ошибке.
Это фундаментальное ограничение архитектуры реализации поддержки AD в PyTorch и причина, по которой мы разработали библиотеку torch.func. Вместо этого используйте эквиваленты API torch.autograd из torch.func:
-
torch.autograd.grad,Tensor.backward->torch.func.vjpилиtorch.func.grad -
torch.autograd.functional.jvp->torch.func.jvp -
torch.autograd.functional.jacobian->torch.func.jacrevилиtorch.func.jacfwd -
torch.autograd.functional.hessian->torch.func.hessian
Ограничения vmap
Примечание
vmap() — наше самое ограничительное преобразование. Преобразования, связанные с grad (grad(), vjp(), jvp()), не имеют этих ограничений. jacfwd() (и hessian(), реализованное с помощью jacfwd()) представляет собой композицию vmap() и jvp(), поэтому на него также распространяются эти ограничения.
vmap(func) — это преобразование, возвращающее функцию, которая применяет func к новым измерениям каждого входного тензора. Удобно представлять vmap как цикл for: для чистых функций (то есть при отсутствии побочных эффектов) vmap(f)(x) эквивалентно следующему:
torch.stack([f(x_i) for x_i in x.unbind(0)])
Мутация: произвольная мутация структур данных Python
При наличии побочных эффектов vmap() перестает работать так, как если бы выполнялся цикл for. Например, следующая функция:
def f(x, list):
list.pop()
print("hello!")
return x.sum(0)
x = torch.randn(3, 1)
lst = [0, 1, 2, 3]
result = vmap(f, in_dims=(0, None))(x, lst)
выведет “hello!” один раз и извлечет только один элемент из lst.
vmap() выполняет f только один раз, поэтому все побочные эффекты возникают только один раз.
Это связано с тем, как реализован vmap. В torch.func есть специальный внутренний класс BatchedTensor. vmap(f)(*inputs) принимает все входные тензоры, преобразует их в BatchedTensor и вызывает f(*batched_tensor_inputs). BatchedTensor переопределяет API PyTorch, обеспечивая пакетное (то есть векторизованное) поведение для каждого оператора PyTorch.
Мутация: операции PyTorch на месте
Возможно, вы попали сюда, получив сообщение об ошибке, связанной с несовместимыми с vmap операциями на месте. vmap() выдаст ошибку при обнаружении неподдерживаемой операции PyTorch на месте; в противном случае выполнение завершится успешно. К неподдерживаемым относятся операции, которые приводят к записи тензора с большим числом элементов в тензор с меньшим числом элементов. Вот пример такой ситуации:
def f(x, y): x.add_(y) return x x = torch.randn(1) y = torch.randn(3, 1) # When vmapped over, looks like it has shape [1] # Raises an error because `x` has fewer elements than `y`. vmap(f, in_dims=(None, 0))(x, y)
x — это тензор с одним элементом, а y — тензор с тремя элементами. x + y содержит три элемента (из-за широковещательной передачи), но попытка записать три элемента обратно в x, содержащий только один элемент, приводит к ошибке, поскольку выполняется попытка записать три элемента в тензор с одним элементом.
Проблемы не возникает, если тензор, в который выполняется запись, является пакетным в рамках vmap() (то есть к нему применяется vmap).
def f(x, y): x.add_(y) return x x = torch.randn(3, 1) y = torch.randn(3, 1) expected = x + y # Does not raise an error because x is being vmapped over. vmap(f, in_dims=(0, 0))(x, y) assert torch.allclose(x, expected)
Один из распространенных способов исправить это — заменить вызовы фабричных функций эквивалентными функциями “new_*”. Например:
- Замените
torch.zeros()наTensor.new_zeros() - Замените
torch.empty()наTensor.new_empty()
Чтобы понять, почему это помогает, рассмотрим следующий пример.
def diag_embed(vec): assert vec.dim() == 1 result = torch.zeros(vec.shape[0], vec.shape[0]) result.diagonal().copy_(vec) return result vecs = torch.tensor([[0., 1, 2], [3., 4, 5]]) # RuntimeError: vmap: inplace arithmetic(self, *extra_args) is not possible ... vmap(diag_embed)(vecs)
Внутри vmap() result — это тензор формы [3, 3]. Однако, хотя кажется, что vec имеет форму [3], фактическая базовая форма vec — [2, 3]. Невозможно скопировать vec в result.diagonal() формы [3], поскольку в нем слишком много элементов.
def diag_embed(vec): assert vec.dim() == 1 result = vec.new_zeros(vec.shape[0], vec.shape[0]) result.diagonal().copy_(vec) return result vecs = torch.tensor([[0., 1, 2], [3., 4, 5]]) vmap(diag_embed)(vecs)
Замена torch.zeros() на Tensor.new_zeros() приводит к тому, что базовый тензор result имеет форму [2, 3, 3]. Теперь можно скопировать vec с базовой формой [2, 3] в result.diagonal().
Мутация: операции PyTorch с out=
vmap() не поддерживает ключевой аргумент out= в операциях PyTorch. При обнаружении такого аргумента в коде будет выдано соответствующее сообщение об ошибке.
Это не фундаментальное ограничение: теоретически мы могли бы поддержать это в будущем, но пока решили этого не делать.
Управляющий поток Python, зависящий от данных
Пока мы не поддерживаем vmap для управляющего потока, зависящего от данных. Управляющий поток зависит от данных, если условие оператора if, цикла while или цикла for — это тензор, к которому применяется vmap. Например, следующий код приведет к ошибке:
def relu(x):
if x > 0:
return x
return 0
x = torch.randn(3)
vmap(relu)(x)
Однако любой управляющий поток, не зависящий от значений в тензорах, к которым применяется vmap, будет работать:
def custom_dot(x):
if x.dim() == 1:
return torch.dot(x, x)
return (x * x).sum()
x = torch.randn(3)
vmap(custom_dot)(x)
JAX поддерживает преобразование управляющего потока, зависящего от данных, с помощью специальных операторов управления потоком (например, jax.lax.cond, jax.lax.while_loop). Мы изучаем возможность добавления аналогичных операторов в PyTorch.
Операции, зависящие от данных (.item())
Мы не поддерживаем и не будем поддерживать vmap для пользовательской функции, вызывающей .item() для тензора. Например, следующий код приведет к ошибке:
def f(x): return x.item() x = torch.randn(3) vmap(f)(x)
Попробуйте переписать код так, чтобы не использовать вызовы .item().
Вы также можете столкнуться с сообщением об ошибке, связанным с использованием .item(), хотя сами его не использовали. В таких случаях PyTorch мог вызвать .item() внутри себя. Сообщите об этой проблеме на GitHub, и мы исправим внутренние механизмы PyTorch.
Операции с динамической формой (nonzero и подобные)
Для vmap(f) необходимо, чтобы применение f к каждому «примеру» во входных данных возвращало тензор одной и той же формы. Такие операции, как torch.nonzero, torch.is_nonzero, не поддерживаются и поэтому приведут к ошибке.
Чтобы понять почему, рассмотрим следующий пример:
xs = torch.tensor([[0, 1, 2], [0, 0, 3]]) vmap(torch.nonzero)(xs)
torch.nonzero(xs[0]) возвращает тензор формы 2, а torch.nonzero(xs[1]) — тензор формы 1. Мы не можем сформировать единый тензор в качестве выходных данных: для этого потребовался бы рваный тензор (а в PyTorch пока нет понятия рваного тензора).
Случайность
Намерения пользователя при вызове операции со случайными данными могут быть неочевидны. В частности, одним пользователям может потребоваться, чтобы случайные значения совпадали между пакетами, а другим — чтобы они различались. Для этого vmap принимает флаг случайности.
Этот флаг можно передать только в vmap; он может принимать одно из трех значений: “error”, “different” или “same”. По умолчанию используется значение error. В режиме “error” любой вызов функции, генерирующей случайные значения, приведет к ошибке с предложением выбрать один из двух других флагов в зависимости от задачи.
При режиме случайности “different” элементы пакета генерируют разные случайные значения. Например:
def add_noise(x): y = torch.randn(()) # y will be different across the batch return x + y x = torch.ones(3) result = vmap(add_noise, randomness="different")(x) # we get 3 different values
При режиме случайности “same” элементы пакета генерируют одинаковые случайные значения. Например:
def add_noise(x): y = torch.randn(()) # y will be the same across the batch return x + y x = torch.ones(3) result = vmap(add_noise, randomness="same")(x) # we get the same value, repeated 3 times
Предупреждение
Наша система определяет поведение случайных значений только для операторов PyTorch и не может управлять поведением других библиотек, например numpy. Это ограничение аналогично ограничениям JAX и предложенным для их решения подходам.
Примечание
Несколько вызовов vmap с любым из поддерживаемых режимов случайности не будут давать одинаковых результатов. Как и в стандартном PyTorch, для воспроизводимости случайных значений можно использовать torch.manual_seed() вне vmap или генераторы.
Примечание
Наконец, наша система случайных значений отличается от JAX, поскольку мы не используем PRNG без состояния, отчасти потому, что PyTorch не поддерживает его в полном объеме. Вместо этого мы ввели систему флагов, позволяющую реализовать наиболее распространенные сценарии использования случайных значений. Если ваш сценарий не относится ни к одному из них, сообщите об этом в виде проблемы.
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/func.ux_limitations.html