PruningContainer
-
class torch.nn.utils.prune.PruningContainer(*args)[source] -
Контейнер, хранящий последовательность методов обрезки для итеративной обрезки. Отслеживает порядок применения методов обрезки и обрабатывает объединение последовательных вызовов обрезки.
Принимает в качестве аргумента экземпляр BasePruningMethod или итерируемый объект из них.
-
add_pruning_method(method)[source] -
Добавляет дочерний метод обрезки
methodв контейнер.- Parameters
-
method (подкласс of 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. Новая частичная маска должна вычисляться для записей или каналов, которые не были обнуленыdefault_mask. Какие части тензораtновая маска будет вычислена, зависит отPRUNING_TYPE(обрабатывается обработчиком типа):- для ‘unstructured’, маска будет вычислена из разобранного списка необрезанных записей;
- для ‘structured’, маска будет вычислена из необрезанных каналов в тензоре;
- для ‘global’, маска будет вычислена по всем записям.
- Parameters
-
-
t (torch.Tensor) – тензор, представляющий параметр для обрезки (той же размерности, что и
default_mask). - default_mask (torch.Tensor) – маска из предыдущей итерации обрезки.
-
t (torch.Tensor) – тензор, представляющий параметр для обрезки (той же размерности, что и
- Returns
-
новая маска, которая объединяет эффекты
default_maskи новой маски из текущей обрезкиmethod(той же размерности, что и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/2.1/generated/torch.nn.utils.prune.PruningContainer.html