torch.nn.utils.prune.global_unstructured
-
torch.nn.utils.prune.global_unstructured(parameters, pruning_method, importance_scores=None, **kwargs)[source] -
Глобально обрезает тензоры, соответствующие всем параметрам в
parametersс применением указанногоpruning_method. Изменяет модули на месте, путём:- добавления именованного буфера, называемого
name+'_mask'соответствующего бинарной маске, применённой к параметруnameметодом обрезки. - замены параметра
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, оно представляет абсолютное число параметров для обрезки.
-
parameters (Iterable of (module, name) tuples) – параметры модели для глобальной обрезки, т.е. путём агрегирования всех весов до того, как будет принято решение, какие из них обрезать. модуль должен быть типа
- 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