Spec-Zone.ru › PyTorch 2

BasePruningMethod

class torch.nn.utils.prune.BasePruningMethod [source]

Абстрактный базовый класс для создания новых методов обрезки.

Предоставляет каркас для настройки, требующий переопределения методов, таких как compute_mask() и apply().

classmethod apply(module, name, *args, importance_scores=None, **kwargs) [source]

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

Parameters
  • module (nn.Module) – модуль, содержащий тензор для обрезки
  • name (str) – имя параметра внутри module, на котором будет выполняться обрезка.
  • args – аргументы, передаваемые подклассу BasePruningMethod
  • importance_scores (torch.Tensor) – тензор оценок важности (той же формы, что и параметр модуля), используемый для вычисления маски обрезки. Значения в этом тензоре указывают на важность соответствующих элементов в параметре, подлежащем обрезке. Если не указано или равно None, параметр будет использован вместо него.
  • kwargs – ключевые аргументы, передаваемые подклассу BasePruningMethod
apply_mask(module) [source]

Просто выполняет умножение параметра, подлежащего обрезке, на сгенерированную маску. Извлекает маску и исходный тензор из модуля и возвращает обрезанную версию тензора.

Parameters

module (nn.Module) – модуль, содержащий тензор для обрезки

Returns

обрезанная версия входного тензора

Return type

pruned_tensor (torch.Tensor)

abstract compute_mask(t, default_mask) [source]

Вычисляет и возвращает маску для входного тензора t. Исходя из базовой default_mask (которая должна быть маской из единиц, если тензор еще не подвергался обрезке), генерирует случайную маску для применения к default_mask в соответствии с рецептом конкретного метода обрезки.

Parameters
  • t (torch.Tensor) – тензор, представляющий оценки важности
  • prune. (параметр для) –
  • default_mask (torch.Tensor) – Базовая маска с предыдущей обрезки
  • iterations –
  • is (которые должны соблюдаться после применения новой маски) –
  • t. (с теми же размерами, что и) –
Returns

маска для применения к t, с теми же размерами, что и t

Return type

mask (torch.Tensor)

prune(t, default_mask=None, importance_scores=None) [source]

Вычисляет и возвращает обрезанную версию входного тензора 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) [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/2.1/generated/torch.nn.utils.prune.BasePruningMethod.html

Spec-Zone.ru

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