BasePruningMethod
-
class torch.nn.utils.prune.BasePruningMethod[source] -
Абстрактный базовый класс для создания новых методов обрезки.
Предоставляет каркас для кастомизации, требующий переопределения методов, таких как
compute_mask()иapply().-
classmethod apply(module, name, *args, importance_scores=None, **kwargs)[source] -
Добавляет предварительный хук forward, который позволяет выполнять обрезку на лету и перепараметризацию тензора относительно исходного тензора и маски обрезки.
- Параметры:
-
- module (nn.Module) – модуль, содержащий тензор для обрезки
-
name (str) – имя параметра в
module, на котором будет действовать обрезка. -
args – аргументы, передаваемые подклассу
BasePruningMethod - importance_scores (torch.Tensor) – тензор оценок важности (той же формы, что и параметр модуля), используемый для вычисления маски обрезки. Значения в этом тензоре указывают на важность соответствующих элементов в параметре, который обрезается. Если не указано или равно None, используется сам параметр.
-
kwargs – ключевые аргументы, передаваемые подклассу
BasePruningMethod
-
apply_mask(module)[source] -
Просто выполняет умножение параметра, подлежащего обрезке, на сгенерированную маску. Извлекает маску и исходный тензор из модуля и возвращает обрезанную версию тензора.
- Параметры:
-
module (nn.Module) – модуль, содержащий тензор для обрезки
- Возвращает:
-
обрезанная версия входного тензора
- Тип возвращаемого значения:
-
pruned_tensor (torch.Tensor)
-
abstract compute_mask(t, default_mask)[source] -
Вычисляет и возвращает маску для входного тензора
t. Начиная с базовойdefault_mask(которая должна быть маской из единиц, если тензор еще не был обрезан), сгенерируйте случайную маску для применения поверхdefault_maskв соответствии со специфической рецептурой метода обрезки.- Параметры:
-
- t (torch.Tensor) – тензор, представляющий оценки важности
- prune. (параметра) –
- default_mask (torch.Tensor) – Базовая маска из предыдущей обрезки
- iterations –
- is (которые необходимо соблюдать после новой маски) –
- t. (применены. Те же размеры, что и) –
- Возвращает:
-
маску для применения к
t, той же размерности, что иt - Тип возвращаемого значения:
-
mask (torch.Tensor)
-
prune(t, default_mask=None, importance_scores=None)[source] -
Вычисляет и возвращает обрезанную версию входного тензора
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)[source] -
Удаляет перепараметризацию обрезки из модуля. Обрезанный параметр с именем
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.BasePruningMethod.html