Spec-Zone.ru › PyTorch 1

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'.
Параметры:
  • module (nn.Module) – модуль, содержащий тензор для обрезки
  • name (str) – имя параметра внутри module на котором будет выполняться обрезка.
  • amount (int или float) – количество параметров для обрезки. Если float, должно быть между 0,0 и 1,0 и представлять долю параметров для обрезки. Если int, оно представляет абсолютное количество параметров для обрезки.
  • dim (int) – индекс измерения, вдоль которого мы определяем каналы для обрезки.
Возвращает:

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

Тип возвращаемого значения:

модуль (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/1.13/generated/torch.nn.utils.prune.random_structured.html

Spec-Zone.ru

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