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