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) – маска предыдущей итерации прореживания.
-
t (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 (torch.Tensor) – тензор для прореживания (с теми же размерами, что и
- Возвращает:
-
прореженная версия тензора
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