Идентичность
-
class torch.nn.utils.prune.Identity[source] -
Утилитарный метод обрезки, который не обрезает никакие единицы, но генерирует параметризацию обрезки с маской из единиц.
-
classmethod apply(module, name)[source] -
Добавляет предварительный хук к методу forward, который позволяет выполнять обрезку на лету и перепараметризовать тензор с точки зрения исходного тензора и маски обрезки.
-
apply_mask(module) -
Просто выполняет умножение параметра, который обрезается, на сгенерированную маску. Извлекает маску и исходный тензор из модуля и возвращает обрезанную версию тензора.
- Параметры:
-
module (nn.Module) – модуль, содержащий тензор для обрезки
- Возвращает:
-
обрезанная версия входного тензора
- Тип возвращаемого значения:
-
pruned_tensor (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'удаляется из буферов.Примечание
Сама обрезка НЕ отменяется и НЕ отменяется!
-
© 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.Identity.html