torch.nn.utils.prune.custom_from_mask
-
torch.nn.utils.prune.custom_from_mask(module, name, mask)[source] -
Обрезает тензор, соответствующий параметру, называемому
nameвmodule, применяя предварительно вычисленную маску вmask. Модифицирует модуль на месте (и также возвращает изменённый модуль) следующим образом:- добавляет именованный буфер, называемый
name+'_mask', соответствующий двоичной маске, применённой к параметруnameметодом обрезки. - заменяет параметр
nameего обрезанной версией, в то время как исходный (необрезанный) параметр сохраняется в новом параметре, названномname+'_orig'.
- Параметры
- Возвращает
-
изменённую (т.е. обрезанную) версию входного модуля
- Тип возвращаемого значения
-
модуль (nn.Module)
Примеры
>>> from torch.nn.utils import prune >>> m = prune.custom_from_mask( ... nn.Linear(5, 3), name='bias', mask=torch.tensor([0, 1, 0]) ... ) >>> print(m.bias_mask) tensor([0., 1., 0.])
- добавляет именованный буфер, называемый
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.nn.utils.prune.custom_from_mask.html