torch.autograd.Function.vmap
-
static Function.vmap(info, in_dims, *args)[исходный код] -
Определяет поведение этого autograd.Function внутри
torch.vmap().Чтобы
torch.autograd.Function()поддерживалtorch.vmap(), необходимо переопределить этот статический метод или присвоитьgenerate_vmap_ruleзначениеTrue(нельзя делать и то и другое).Если вы решили переопределить этот статический метод, он должен принимать
- объект
infoв качестве первого аргумента.info.batch_sizeзадаёт размер измерения, по которому выполняется vmap, аinfo.randomness— параметр случайности, передаваемый вtorch.vmap(). - кортеж
in_dimsв качестве второго аргумента. Для каждого аргумента вargsвin_dimsимеется соответствующийOptional[int]. ЭтоNone, если аргумент не является тензором или если для него не выполняется vmap; в противном случае это целое число, указывающее, по какому измерению тензора выполняется vmap. -
*args, совпадающий с аргументамиforward().
Статический метод vmap возвращает кортеж из
(output, out_dims). Как иin_dims,out_dimsдолжен иметь ту же структуру, что иoutput, и содержать по одномуout_dimдля каждого результата, указывающему, содержит ли результат измерение vmap и по какому индексу оно расположено.Подробнее см. в разделе Расширение torch.func с помощью autograd.Function.
- объект
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.autograd.Function.vmap.html