PruningContainer
-
class torch.nn.utils.prune.PruningContainer(*args)[source] -
Контейнер, хранящий последовательность методов обрезки для итеративной обрезки. Отслеживает порядок применения методов обрезки и обрабатывает объединение последовательных вызовов обрезки.
Принимает в качестве аргумента экземпляр BasePruningMethod или итерируемый объект таких экземпляров.
-
add_pruning_method(method)[source] -
Добавляет дочерний метод обрезки
methodв контейнер.- Parameters:
-
method (подкласс BasePruningMethod) – дочерний метод обрезки, который необходимо добавить в контейнер.
-
classmethod apply(module, name, *args, importance_scores=None, **kwargs) -
Добавляет предварительный хук forward, который позволяет выполнять обрезку на лету и перепараметризовать тензор в терминах исходного тензора и маски обрезки.
- Parameters:
-
- module (nn.Module) – модуль, содержащий тензор для обрезки
-
name (str) – имя параметра внутри
module, на котором будет действовать обрезка. -
args – аргументы, передаваемые подклассу
BasePruningMethod - importance_scores (torch.Tensor) – тензор оценок важности (такой же формы, как параметр модуля), используемый для вычисления маски обрезки. Значения в этом тензоре указывают на важность соответствующих элементов в параметре, который обрезается. Если не указано или равно None, используется сам параметр.
-
kwargs – ключевые аргументы, передаваемые подклассу
BasePruningMethod
-
apply_mask(module) -
Просто обрабатывает умножение параметра, который обрезается, на сгенерированную маску. Извлекает маску и исходный тензор из модуля и возвращает обрезанную версию тензора.
- Parameters:
-
module (nn.Module) – модуль, содержащий тензор для обрезки
- Returns:
-
обрезанная версия входного тензора
- Return type:
-
pruned_tensor (torch.Tensor)
-
compute_mask(t, default_mask)[source] -
Применяет последний
method, вычисляя новые частичные маски и возвращая их комбинацию сdefault_mask. Новая частичная маска должна быть вычислена для элементов или каналов, которые не были обнулены предыдущей обрезкой. Какие части тензораtбудут использованы для вычисления новой маски зависит отPRUNING_TYPE(обрабатывается обработчиком типа):- для ‘unstructured’, маска будет вычислена из «раскрученной» последовательности необрезанных элементов;
- для ‘structured’, маска будет вычислена из необрезанных каналов в тензоре;
- для ‘global’, маска будет вычислена по всем элементам.
- Parameters:
-
-
t (torch.Tensor) – тензор, представляющий параметр для обрезки (той же размерности, что и
default_mask). - default_mask (torch.Tensor) – маска из предыдущей итерации обрезки.
-
t (torch.Tensor) – тензор, представляющий параметр для обрезки (той же размерности, что и
- Returns:
-
новая маска, которая объединяет эффекты
default_maskи новой маски от текущей операции обрезки (той же размерности, что иdefault_maskиt). - Return type:
-
mask (torch.Tensor)
-
prune(t, default_mask=None, importance_scores=None) -
Вычисляет и возвращает обрезанную версию входного тензора
tв соответствии с правилом обрезки, указанным вcompute_mask().- Parameters:
-
-
t (torch.Tensor) – тензор для обрезки (той же размерности, что и
default_mask). -
importance_scores (torch.Tensor) – тензор оценок важности (той же формы, что и
t) для вычисления маски обрезкиt. Значения в этом тензоре указывают на важность соответствующих элементов вtпараметре, который обрезается. Если не указано или равно None, используется тензорt. - default_mask (torch.Tensor, optional) – маска из предыдущей итерации обрезки, если есть. Используется для определения части тензора, на которой должна действовать обрезка. Если None, используется маска из единиц.
-
t (torch.Tensor) – тензор для обрезки (той же размерности, что и
- Returns:
-
обрезанная версия тензора
t.
-
remove(module) -
Удаляет перепараметризацию обрезки из модуля. Обрезанный параметр с именем
nameостается постоянно обрезанным, а параметр с именемname+'_orig'удаляется из списка параметров. Аналогично, буфер с именемname+'_mask'удаляется из буферов.Примечание
Сама обрезка НЕ отменяется или не восстанавливается!
-
© 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.PruningContainer.html