Spec-Zone.ru › PyTorch 2.14

PruningContainer

class torch.nn.utils.prune.PruningContainer(*args) [исходный код]

Контейнер, содержащий последовательность методов прореживания для итеративного прореживания.

Отслеживает порядок применения методов прореживания и обрабатывает объединение последовательных вызовов прореживания.

В качестве аргумента принимает экземпляр BasePruningMethod или итерируемый объект с такими экземплярами.

add_pruning_method(method) [исходный код]

Добавляет дочерний method прореживания в контейнер.

Параметры:

method (подкласс BasePruningMethod) – дочерний метод прореживания, добавляемый в контейнер.

classmethod apply(module, name, *args, importance_scores=None, **kwargs) [исходный код]

Выполняет прореживание на лету и репараметризацию тензора.

Добавляет предварительный хук прямого прохода, который позволяет выполнять прореживание на лету, и репараметризует тензор через исходный тензор и маску прореживания.

Параметры:
  • module (nn.Module) – модуль, содержащий тензор для прореживания
  • name (str) – имя параметра в module, к которому будет применено прореживание.
  • args – аргументы, передаваемые подклассу BasePruningMethod
  • importance_scores (torch.Tensor) – тензор оценок важности (той же формы, что и параметр модуля), используемый для вычисления маски прореживания. Значения в этом тензоре указывают важность соответствующих элементов прореживаемого параметра. Если значение не задано или равно None, вместо него будет использоваться параметр.
  • kwargs – именованные аргументы, передаваемые подклассу BasePruningMethod
apply_mask(module) [исходный код]

Выполняет умножение прореживаемого параметра на созданную маску.

Извлекает маску и исходный тензор из модуля и возвращает прореженную версию тензора.

Параметры:

module (nn.Module) – модуль, содержащий тензор для прореживания

Возвращает:

прореженную версию входного тензора

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

pruned_tensor (torch.Tensor)

compute_mask(t, default_mask) [исходный код]

Применяет последний method, вычисляя новые частичные маски и возвращая их объединение с default_mask.

Новая частичная маска должна вычисляться для элементов или каналов, которые не были обнулены с помощью default_mask. От PRUNING_TYPE (обрабатываемого обработчиком типа) зависит, на каких частях тензора t будет вычислена новая маска:

  • для ‘unstructured’ маска будет вычисляться по развёрнутому списку немаскированных элементов;
  • для ‘structured’ маска будет вычисляться по немаскированным каналам тензора;
  • для ‘global’ маска будет вычисляться для всех элементов.
Параметры:
  • t (torch.Tensor) – тензор, представляющий прореживаемый параметр (с теми же размерами, что и default_mask).
  • default_mask (torch.Tensor) – маска предыдущей итерации прореживания.
Возвращает:

новая маска, объединяющая эффекты default_mask и новой маски текущего method прореживания (с теми же размерами, что и default_mask и t).

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

mask (torch.Tensor)

prune(t, default_mask=None, importance_scores=None) [исходный код]

Вычисляет и возвращает прореженную версию входного тензора t.

В соответствии с правилом прореживания, заданным в compute_mask().

Параметры:
  • t (torch.Tensor) – тензор для прореживания (с теми же размерами, что и default_mask).
  • importance_scores (torch.Tensor) – тензор оценок важности (той же формы, что и t), используемый для вычисления маски прореживания t. Значения в этом тензоре указывают важность соответствующих элементов t, который прореживается. Если значение не задано или равно None, вместо него будет использоваться тензор t.
  • default_mask (torch.Tensor, необязательный) – маска предыдущей итерации прореживания, если она есть. Учитывается при определении части тензора, к которой следует применить прореживание. Если значение равно None, по умолчанию используется маска из единиц.
Возвращает:

прореженная версия тензора t.

remove(module) [исходный код]

Удаляет репараметризацию прореживания из модуля.

Прореженный параметр с именем name остаётся прореженным без возможности отмены, а параметр с именем name+'_orig' удаляется из списка параметров. Аналогично, буфер с именем name+'_mask' удаляется из списка буферов.

Примечание

Само прореживание НЕ отменяется и не обращается!

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

Spec-Zone.ru

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