Spec-Zone.ru › PyTorch 2

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

Spec-Zone.ru

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