Spec-Zone.ru › PyTorch 1

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]

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

Параметры:
  • 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/1.13/generated/torch.nn.utils.prune.LnStructured.html

Spec-Zone.ru

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