Spec-Zone.ru › PyTorch 2

LnStructured

class torch.nn.utils.prune.LnStructured(amount, n, dim=-1) [source]

Обрезает целые (в настоящее время не обрезанные) каналы в тензоре на основе их Ln-нормы.

Параметры
  • amount (int или float) – количество каналов для обрезки. Если float, должно быть между 0,0 и 1,0 и представлять долю параметров для обрезки. Если int, оно представляет собой абсолютное количество параметров для обрезки.
  • n (int, float, inf, -inf, 'fro', 'nuc') – См. документацию по допустимым значениям для аргумента p в torch.norm().
  • dim (int, необязательно) – индекс измерения, по которому определяются каналы для обрезки. По умолчанию: -1.
classmethod apply(module, name, amount, n, dim, importance_scores=None) [source]

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

Параметры
  • module (nn.Module) – модуль, содержащий тензор для обрезки
  • name (str) – имя параметра внутри module, на котором будет действовать обрезка.
  • amount (int или float) – количество параметров для обрезки. Если float, должно быть между 0,0 и 1,0 и представлять долю параметров для обрезки. Если int, оно представляет собой абсолютное количество параметров для обрезки.
  • n (int, float, inf, -inf, 'fro', 'nuc') – См. документацию по допустимым значениям для аргумента p в torch.norm().
  • dim (int) – индекс измерения, по которому определяются каналы для обрезки.
  • importance_scores (torch.Tensor) – тензор оценок важности (той же формы, что и параметр модуля), используемый для вычисления маски обрезки. Значения в этом тензоре указывают на важность соответствующих элементов в параметре, подлежащем обрезке. Если не указано или None, используется параметр модуля.
apply_mask(module)

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

Параметры

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

Возвращает

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

Тип возвращаемого значения

pruned_tensor (torch.Tensor)

compute_mask(t, default_mask) [source]

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

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

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

Тип возвращаемого значения

mask (torch.Tensor)

Возбуждает

IndexError – если self.dim >= len(t.shape)

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/2.1/generated/torch.nn.utils.prune.LnStructured.html

Spec-Zone.ru

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