torch.nn.utils.prune.random_structured
-
torch.nn.utils.prune.random_structured(module, name, amount, dim)[source] -
Удаляет тензор, соответствующий параметру, названному
nameвmodule, удаляя указанноеamount(в настоящее время не обрезаных) каналов вдоль указанногоdimслучайным образом. Изменяет модуль на месте (а также возвращает изменённый модуль) следующим образом:- добавляет именованный буфер, называемый
name+'_mask', соответствующий двоичной маске, применённой к параметруnameметодом обрезки. - заменяет параметр
nameего обрезанной версией, а исходный (не обрезанный) параметр сохраняется в новом параметре под именемname+'_orig'.
- Параметры:
-
- module (nn.Module) – модуль, содержащий тензор для обрезки
-
name (str) – имя параметра внутри
moduleна котором будет выполняться обрезка. -
amount (int или float) – количество параметров для обрезки. Если
float, должно быть между 0,0 и 1,0 и представлять долю параметров для обрезки. Еслиint, оно представляет абсолютное количество параметров для обрезки. - dim (int) – индекс измерения, вдоль которого мы определяем каналы для обрезки.
- Возвращает:
-
изменённый (т. е. обрезанный) вариант входного модуля
- Тип возвращаемого значения:
-
модуль (nn.Module)
Примеры
>>> m = prune.random_structured( ... nn.Linear(5, 3), 'weight', amount=3, dim=1 ... ) >>> columns_pruned = int(sum(torch.sum(m.weight, dim=0) == 0)) >>> print(columns_pruned) 3
- добавляет именованный буфер, называемый
© 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.random_structured.html