Spec-Zone.ru › PyTorch 1

CustomFromMask

class torch.nn.utils.prune.CustomFromMask(mask) [source]
classmethod apply(module, name, mask) [source]

Добавляет предварительный хук для выполнения обрезки на лету и перепараметризации тензора относительно исходного тензора и маски обрезки.

Параметры:
  • module (nn.Module) – модуль, содержащий тензор для обрезки
  • name (строка) – имя параметра внутри module, на котором будет выполняться обрезка.
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.

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.CustomFromMask.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API