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. - Тип возвращаемого значения
Один из примеров использования
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