torch.nn.utils.prune.custom_from_mask
-
torch.nn.utils.prune.custom_from_mask(module, name, mask)[исходный код] -
Обрезает тензор, соответствующий параметру с именем
nameвmodule, применяя предварительно вычисленную маску изmask.Изменяет модуль на месте (и также возвращает изменённый модуль), выполняя следующие действия:
- добавляет именованный буфер с именем
name+'_mask', соответствующий двоичной маске, применённой методом обрезки к параметруname. - заменяет параметр
nameего обрезанной версией, сохраняя исходный (необрезанный) параметр в новом параметре с именемname+'_orig'.
- Параметры:
- Возвращает:
-
изменённую (то есть обрезанную) версию входного модуля
- Тип возвращаемого значения:
-
module (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.])
- добавляет именованный буфер с именем
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.nn.utils.prune.custom_from_mask.html