Spec-Zone.ru › PyTorch 2

torch.nn.utils.prune.global_unstructured

torch.nn.utils.prune.global_unstructured(parameters, pruning_method, importance_scores=None, **kwargs) [source]

Глобально обрезает тензоры, соответствующие всем параметрам в parameters с применением указанного pruning_method. Изменяет модули на месте, путём:

  1. добавления именованного буфера, называемого name+'_mask' соответствующего бинарной маске, применённой к параметру name методом обрезки.
  2. замены параметра name его обрезанной версией, при этом исходный (необрезной) параметр хранится в новом параметре под именем name+'_orig'.
Parameters
  • parameters (Iterable of (module, name) tuples) – параметры модели для глобальной обрезки, т.е. путём агрегирования всех весов до того, как будет принято решение, какие из них обрезать. модуль должен быть типа nn.Module, а имя — строкой.
  • pruning_method (function) – допустимая функция обрезки из этого модуля или пользовательская, реализованная пользователем, удовлетворяющая руководству по реализации и имеющая PRUNING_TYPE='unstructured'.
  • importance_scores (dict) – словарь, сопоставляющий кортежи (модуль, имя) соответствующим тензорам оценок важности параметра. Тензор должен иметь ту же форму, что и параметр, и используется для вычисления маски обрезки. Если не указано или равно None, параметр будет использован вместо оценок его важности.
  • kwargs – другие ключевые аргументы, такие как: amount (int или float): количество параметров для обрезки по указанным параметрам. Если float, должно быть между 0,0 и 1,0 и представлять долю параметров для обрезки. Если int, оно представляет абсолютное число параметров для обрезки.
Raises

TypeError – если PRUNING_TYPE != 'unstructured'

Примечание

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

Примеры

>>> from torch.nn.utils import prune
>>> from collections import OrderedDict
>>> net = nn.Sequential(OrderedDict([
...     ('first', nn.Linear(10, 4)),
...     ('second', nn.Linear(4, 1)),
... ]))
>>> parameters_to_prune = (
...     (net.first, 'weight'),
...     (net.second, 'weight'),
... )
>>> prune.global_unstructured(
...     parameters_to_prune,
...     pruning_method=prune.L1Unstructured,
...     amount=10,
... )
>>> print(sum(torch.nn.utils.parameters_to_vector(net.buffers()) == 0))
tensor(10)

© 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.global_unstructured.html

Spec-Zone.ru

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