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