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)
Примеры
>>> 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/1.13/generated/torch.nn.utils.prune.custom_from_mask.html