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