Spec-Zone.ru › PyTorch 2

torch.vmap

torch.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 также можно использовать для вычисления пакетных градиентов при составлении с автоградом.

Примечание

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, попробуйте использовать ненулевое значение для chunk_size.
Возвращаемое значение

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

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

Callable

Один пример использования vmap() — вычисление пакетных скалярных произведений. PyTorch не предоставляет пакетный torch.dot API; вместо безуспешного поиска в документации, используйте 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 вычисляет vjps (векторно-якобианские произведения). Вычисление полной матрицы Якоби для некоторой функции 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.vmap.html

Spec-Zone.ru

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