torch.nn.utils.prune.ln_structured
-
torch.nn.utils.prune.ln_structured(module, name, amount, n, dim, importance_scores=None)[source] -
Удаляет тензор, соответствующий параметру, называемому
nameвmodule, удаляя указаннуюamount(в настоящее время не обрезанных) каналов вдоль указанногоdimс наименьшей Ln-нормой. Изменяет модуль на месте (а также возвращает измененный модуль), выполняя:- добавление именованного буфера, называемого
name+'_mask', соответствующего двоичной маске, применённой к параметруnameметодом обрезки. - замена параметра
nameего обрезанной версией, а оригинальный (не обрезанный) параметр хранится в новом параметре, названномname+'_orig'.
- Parameters:
-
- module (nn.Module) – модуль, содержащий тензор для обрезки
-
name (str) – имя параметра в
module, на котором будет действовать обрезка. -
amount (int or float) – количество параметров для обрезки. Если
float, должно быть между 0,0 и 1,0 и представлять долю параметров для обрезки. Еслиint, оно представляет собой абсолютное число параметров для обрезки. -
n (int, float, inf, -inf, 'fro', 'nuc') – Смотрите документацию допустимых значений для аргумента
pвtorch.norm(). - dim (int) – индекс измерения, вдоль которого мы определяем каналы для обрезки.
- importance_scores (torch.Tensor) – тензор оценок важности (той же формы, что и параметр модуля), используемый для вычисления маски для обрезки. Значения в этом тензоре указывают на важность соответствующих элементов в параметре, подлежащем обрезке. Если не указано или равно None, используется параметр модуля.
- Returns:
-
измененный (т.е. обрезенный) вариант входного модуля
- Return type:
-
модуль (nn.Module)
Примеры
>>> m = prune.ln_structured( ... nn.Conv2d(5, 3, 2), 'weight', amount=0.3, dim=1, n=float('-inf') ... ) - добавление именованного буфера, называемого
© 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.ln_structured.html