Spec-Zone.ru › PyTorch 1

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 (Итерируемый список кортежей (модуль, имя)) – параметры модели для глобальной обрезки, т.е. путем агрегирования всех весов перед принятием решения о том, какие из них следует обрезать. Модуль должен быть типа nn.Module, а имя должно быть строкой.
  • pruning_method (функция) – допустимая функция обрезки из этого модуля или пользовательская функция, реализованная пользователем, которая удовлетворяет принципам реализации и имеет PRUNING_TYPE='unstructured'.
  • importance_scores (словарь) – словарь, сопоставляющий кортежи (модуль, имя) соответствующему тензору оценок важности параметра. Тензор должен иметь тот же размер, что и параметр, и используется для вычисления маски обрезки. Если не указано или равно None, параметр будет использован вместо своих оценок важности.
  • kwargs – другие ключевые аргументы, такие как: amount (int или float): количество параметров для обрезки по указанным параметрам. Если float, должно быть между 0,0 и 1,0 и представлять собой долю параметров для обрезки. Если int, оно представляет собой абсолютное количество параметров для обрезки.
Возбуждает:

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

Примечание

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

Примеры

>>> 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, dtype=torch.uint8)

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

Spec-Zone.ru

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