Spec-Zone.ru › PyTorch 2

torch.func.vmap

torch.func.vmap(func, in_dims=0, out_dims=0, randomness='error', *, chunk_size=None)

vmap — это функция векторного отображения; vmap(func) возвращает новую функцию, которая отображает func по некоторому измерению входных данных. Семантически, vmap помещает отображение в операции PyTorch, вызываемые func, эффективно векторизируя эти операции.

vmap полезна для обработки размерностей пакетных данных: можно написать функцию func, которая работает с примерами, а затем поднять её до функции, которая может принимать пакеты примеров с vmap(func). vmap также может использоваться для вычисления пакетных градиентов при композиции с autograd.

Примечание

torch.vmap() алиасирована с torch.func.vmap() для удобства. Используйте ту, которая вам нужна.

Параметры
  • func (функция) – Функция Python, которая принимает один или несколько аргументов. Должна возвращать один или несколько тензоров.
  • in_dims (int или вложенная структура) – Указывает, какое измерение входных данных должно быть отображено. in_dims должна иметь структуру, подобную входным данным. Если in_dim для конкретного входного значения равно None, это означает, что измерения отображения нет. По умолчанию: 0.
  • out_dims (int или кортеж[int]) – Указывает, где измерение отображения должно появиться в выходных данных. Если out_dims является кортежем, то он должен содержать по одному элементу для каждого выходного значения. По умолчанию: 0.
  • randomness (str) – Указывает, одинаковая или разная случайность должна быть в этом vmap для разных пакетов. Если ‘different’, случайность для каждого пакета будет разной. Если ‘same’, случайность будет одинаковой для всех пакетов. Если ‘error’, любые вызовы случайных функций вызовут ошибку. По умолчанию: ‘error’. ПРЕДУПРЕЖДЕНИЕ: этот флаг применим только к случайным операциям PyTorch и не применим к модулю случайных чисел Python или случайности numpy.
  • chunk_size (None или int) – Если None (по умолчанию), применяется одно vmap к входным данным. Если не None, то вычисляется vmap по chunk_size образцов за раз. Обратите внимание, что chunk_size=1 эквивалентно вычислению vmap с помощью цикла for. Если у вас возникают проблемы с памятью при вычислении vmap, попробуйте задать не None значение chunk_size.
Возвращаемое значение

Возвращает новую «пакетную» функцию. Она принимает те же входные данные, что и func, за исключением того, что каждый вход имеет дополнительное измерение в указанном индексе in_dims. Она возвращает те же выходные данные, что и func, за исключением того, что каждое выходное значение имеет дополнительное измерение в указанном индексе out_dims.

Тип возвращаемого значения

Callable

Один из примеров использования vmap() — вычисление пакетных скалярных произведений. PyTorch не предоставляет API для пакетных torch.dot; вместо того, чтобы безуспешно искать в документации, используйте vmap() для создания новой функции.

>>> torch.dot                            # [D], [D] -> []
>>> batched_dot = torch.func.vmap(torch.dot)  # [N, D], [N, D] -> [N]
>>> x, y = torch.randn(2, 5), torch.randn(2, 5)
>>> batched_dot(x, y)

vmap() может быть полезна для скрытия размерностей пакетов, что упрощает работу с моделью.

>>> batch_size, feature_size = 3, 5
>>> weights = torch.randn(feature_size, requires_grad=True)
>>>
>>> def model(feature_vec):
>>>     # Very simple linear model with activation
>>>     return feature_vec.dot(weights).relu()
>>>
>>> examples = torch.randn(batch_size, feature_size)
>>> result = torch.vmap(model)(examples)

vmap() также может помочь векторизовать вычисления, которые ранее были трудно или невозможно выполнить в пакетном режиме. Один из примеров — вычисление градиентов высшего порядка. Движок PyTorch autograd вычисляет векторно-якобианские произведения (VJП). Вычисление полной матрицы Якоби для некоторой функции f: R^N -> R^N обычно требует N вызовов autograd.grad, по одному на строку Якоби. Используя vmap(), мы можем векторизовать всё вычисление, вычисляя Якобиан в одном вызове autograd.grad.

>>> # Setup
>>> N = 5
>>> f = lambda x: x ** 2
>>> x = torch.randn(N, requires_grad=True)
>>> y = f(x)
>>> I_N = torch.eye(N)
>>>
>>> # Sequential approach
>>> jacobian_rows = [torch.autograd.grad(y, x, v, retain_graph=True)[0]
>>>                  for v in I_N.unbind()]
>>> jacobian = torch.stack(jacobian_rows)
>>>
>>> # vectorized gradient computation
>>> def get_vjp(v):
>>>     return torch.autograd.grad(y, x, v)
>>> jacobian = torch.vmap(get_vjp)(I_N)

vmap() также может быть вложенной, создавая выходные данные с несколькими пакетными измерениями

>>> torch.dot                            # [D], [D] -> []
>>> batched_dot = torch.vmap(torch.vmap(torch.dot))  # [N1, N0, D], [N1, N0, D] -> [N1, N0]
>>> x, y = torch.randn(2, 3, 5), torch.randn(2, 3, 5)
>>> batched_dot(x, y) # tensor of size [2, 3]

Если входы не сгруппированы по первому измерению, in_dims указывает измерение, по которому каждый вход сгруппирован, как

>>> torch.dot                            # [N], [N] -> []
>>> batched_dot = torch.vmap(torch.dot, in_dims=1)  # [N, D], [N, D] -> [D]
>>> x, y = torch.randn(2, 5), torch.randn(2, 5)
>>> batched_dot(x, y)   # output is [5] instead of [2] if batched along the 0th dimension

Если имеется несколько входов, каждый из которых сгруппирован по различным измерениям, in_dims должен быть кортежем с размерностью пакета для каждого входа, как

>>> torch.dot                            # [D], [D] -> []
>>> batched_dot = torch.vmap(torch.dot, in_dims=(0, None))  # [N, D], [D] -> [N]
>>> x, y = torch.randn(2, 5), torch.randn(5)
>>> batched_dot(x, y) # second arg doesn't have a batch dim because in_dim[1] was None

Если вход представляет собой Python-структуру, in_dims должен быть кортежем, содержащим структуру, соответствующую форме входа:

>>> f = lambda dict: torch.dot(dict['x'], dict['y'])
>>> x, y = torch.randn(2, 5), torch.randn(5)
>>> input = {'x': x, 'y': y}
>>> batched_dot = torch.vmap(f, in_dims=({'x': 0, 'y': None},))
>>> batched_dot(input)

По умолчанию выходные данные сгруппированы по первому измерению. Однако их можно сгруппировать по любому измерению, используя out_dims

>>> f = lambda x: x ** 2
>>> x = torch.randn(2, 5)
>>> batched_pow = torch.vmap(f, out_dims=1)
>>> batched_pow(x) # [5, 2]

Для любой функции, использующей kwargs, возвращаемая функция не будет группировать kwargs, а будет принимать kwargs

>>> x = torch.randn([2, 5])
>>> def fn(x, scale=4.):
>>>   return x * scale
>>>
>>> batched_pow = torch.vmap(fn)
>>> assert torch.allclose(batched_pow(x), x * 4)
>>> batched_pow(x, scale=x) # scale is not batched, output has shape [2, 2, 5]

Примечание

vmap не предоставляет общую автогруппировку или обработку последовательностей переменной длины «из коробки».

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

Spec-Zone.ru

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