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