Spec-Zone.ru › PyTorch 2

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) – маска из предыдущей итерации обрезки.
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, по умолчанию используется маска из единиц.
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

Spec-Zone.ru

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