torch.autograd.Function.vmap
-
static Function.vmap(info, in_dims, *args)[source] -
Определяет правило поведения этого 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.
- объект
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.autograd.Function.vmap.html