Spec-Zone.ru › PyTorch 2

torch.nn.utils.prune.random_structured

torch.nn.utils.prune.random_structured(module, name, amount, dim) [source]

Удаляет тензор, соответствующий параметру под названием name в module, удаляя указанную amount (в настоящее время не удалённых) каналов вдоль указанного dim выбранного случайным образом. Изменяет модуль на месте (а также возвращает изменённый модуль), путём:

  1. добавления именованного буфера под названием name+'_mask', соответствующего двоичной маске, применяемой к параметру name методом обрезки.
  2. замены параметра name его обрезанной версией, а оригинальный (не обрезанный) параметр хранится в новом параметре с именем name+'_orig'.
Parameters
  • module (nn.Module) – модуль, содержащий тензор для обрезки
  • name (str) – имя параметра внутри module, к которому будет применена обрезка.
  • amount (int or float) – количество параметров для обрезки. Если float, должно быть в диапазоне от 0,0 до 1,0 и представлять собой долю параметров для обрезки. Если int, оно представляет собой абсолютное число параметров для обрезки.
  • dim (int) – индекс измерения, по которому мы определяем каналы для обрезки.
Returns

изменённый (т.е. обрезанный) вариант входного модуля

Return type

модуль (nn.Module)

Примеры

>>> m = prune.random_structured(
...     nn.Linear(5, 3), 'weight', amount=3, dim=1
... )
>>> columns_pruned = int(sum(torch.sum(m.weight, dim=0) == 0))
>>> print(columns_pruned)
3

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.nn.utils.prune.random_structured.html

Spec-Zone.ru

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