Spec-Zone.ru › PyTorch 2

Ограничения UX

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()), преобразование, возможно, не сможет преобразовать её. В случае невозможности преобразования вы получите сообщение об ошибке.

Это фундаментальное ограничение дизайна в том, как реализована поддержка AD PyTorch, и причина, по которой мы разработали библиотеку torch.func. Пожалуйста, используйте вместо этого эквиваленты torch.func API torch.autograd: - 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(), 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) принимает все входные тензоры, преобразует их в BatchedTensors и вызывает 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().

Мутация: out= Операции PyTorch

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 и он может принимать три значения: «ошибка», «разные» или «одинаковые», по умолчанию установлено значение «ошибка». В режиме «ошибка» любой вызов случайной функции вызовет ошибку с просьбой использовать один из других двух флагов, в зависимости от конкретного случая использования.

При «различной» случайности элементы в пачке генерируют разные случайные значения. Например,

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

При «одинаковой» случайности элементы в пачке генерируют одинаковые случайные значения. Например,

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 не имеет полной поддержки бессостоятельного PRNG. Вместо этого мы ввели систему флагов, чтобы позволить наиболее распространённые формы случайности, которые мы видим. Если ваш случай использования не соответствует этим формам случайности, пожалуйста, создайте вопрос.

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

Spec-Zone.ru

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