Spec-Zone.ru › PyTorch 2

L1Неструктурированный

class torch.nn.utils.prune.L1Unstructured(amount) [source]

Обрезает (в настоящее время не обрезанные) блоки в тензоре, обнуляя те, у которых норма L1 минимальна.

Параметры

amount (int или float) – количество параметров для обрезки. Если float, должно быть от 0.0 до 1.0 и представлять долю параметров для обрезки. Если int, оно представляет собой абсолютное число параметров для обрезки.

classmethod apply(module, name, amount, importance_scores=None) [source]

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

Параметры
  • module (nn.Module) – модуль, содержащий тензор для обрезки
  • name (str) – имя параметра внутри module, на котором будет действовать обрезка.
  • amount (int или float) – количество параметров для обрезки. Если float, должно быть от 0.0 до 1.0 и представлять долю параметров для обрезки. Если int, оно представляет собой абсолютное число параметров для обрезки.
  • importance_scores (torch.Tensor) – тензор оценок важности (той же формы, что и параметр модуля), используемый для вычисления маски обрезки. Значения в этом тензоре указывают на важность соответствующих элементов в параметре, подлежащем обрезке. Если не указано или равно None, используется параметр модуля.
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/2.1/generated/torch.nn.utils.prune.L1Unstructured.html

Spec-Zone.ru

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