LnStructured
-
class torch.nn.utils.prune.LnStructured(amount, n, dim=- 1)[source] -
Обрезка целых (в данный момент не обрезанных) каналов в тензоре на основе их L
n-нормы.- Параметры:
-
-
amount (int или float) – количество каналов для обрезки. Если
float, должно быть в диапазоне от 0,0 до 1,0 и представлять долю параметров для обрезки. Еслиint, оно представляет собой абсолютное количество параметров для обрезки. -
n (int, float, inf, -inf, 'fro', 'nuc') – См. документацию допустимых значений для аргумента
pвtorch.norm(). - dim (int, необязательно) – индекс измерения, по которому определяются каналы для обрезки. По умолчанию: -1.
-
amount (int или float) – количество каналов для обрезки. Если
-
classmethod apply(module, name, amount, n, dim, importance_scores=None)[source] -
Добавляет предварительный хук forward, который позволяет производить обрезку на лету и перепараметризовать тензор относительно исходного тензора и маски обрезки.
- Параметры:
-
- module (nn.Module) – модуль, содержащий тензор для обрезки
-
name (str) – имя параметра в
module, на котором будет действовать обрезка. -
amount (int или float) – количество параметров для обрезки. Если
float, должно быть в диапазоне от 0,0 до 1,0 и представлять долю параметров для обрезки. Еслиint, оно представляет собой абсолютное количество параметров для обрезки. -
n (int, float, inf, -inf, 'fro', 'nuc') – См. документацию допустимых значений для аргумента
pвtorch.norm(). - dim (int) – индекс измерения, по которому определяются каналы для обрезки.
- importance_scores (torch.Tensor) – тензор оценок важности (такой же формы, как параметр модуля), используемый для вычисления маски обрезки. Значения в этом тензоре указывают на важность соответствующих элементов в параметре, подлежащем обрезке. Если не указан или None, используется параметр модуля.
-
apply_mask(module) -
Просто выполняет умножение между параметром, подлежащим обрезке, и сгенерированной маской. Извлекает маску и исходный тензор из модуля и возвращает обрезанную версию тензора.
- Параметры:
-
module (nn.Module) – модуль, содержащий тензор для обрезки
- Возвращает:
-
обрезанная версия входного тензора
- Тип возвращаемого значения:
-
pruned_tensor (torch.Tensor)
-
compute_mask(t, default_mask)[source] -
Вычисляет и возвращает маску для входного тензора
t. Начиная с базовойdefault_mask(которая должна быть маской из единиц, если тензор еще не обрезался), генерирует маску для применения поверхdefault_maskпутем обнуления каналов вдоль указанного измерения с наименьшей Ln-нормой.- Параметры:
-
- t (torch.Tensor) – тензор, представляющий параметр для обрезки
-
default_mask (torch.Tensor) – Базовая маска из предыдущих итераций обрезки, которая должна быть учтена после применения новой маски. Та же размерность, что и у
t.
- Возвращает:
-
маска для применения к
t, той же размерности, что и уt - Тип возвращаемого значения:
-
mask (torch.Tensor)
- Возможные исключения:
-
IndexError – если
self.dim >= len(t.shape)
-
prune(t, default_mask=None, importance_scores=None) -
Вычисляет и возвращает обрезанную версию входного тензора
tв соответствии с правилом обрезки, указанным вcompute_mask().- Параметры:
-
-
t (torch.Tensor) – тензор для обрезки (той же размерности, что и
default_mask). -
importance_scores (torch.Tensor) – тензор оценок важности (той же формы, что и
t) используемый для вычисления маски обрезкиt. Значения в этом тензоре указывают на важность соответствующих элементов вt, который подлежит обрезке. Если не указан или None, используется тензорt. - default_mask (torch.Tensor, необязательно) – маска с предыдущей итерации обрезки, если она есть. Она учитывается при определении, какая часть тензора подвергается обрезке. Если None, по умолчанию используется маска из единиц.
-
t (torch.Tensor) – тензор для обрезки (той же размерности, что и
- Возвращает:
-
обрезанная версия тензора
t.
-
remove(module) -
Удаляет перепараметризацию обрезки из модуля. Обрезанный параметр с именем
nameостается постоянно обрезанным, а параметр с именемname+'_orig'удаляется из списка параметров. Аналогично, буфер с именемname+'_mask'удаляется из буферов.Примечание
Сама обрезка НЕ отменяется и НЕ отменяется!
© 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.LnStructured.html