Spec-Zone.ru › PyTorch 2.14

torch.nn.utils.prune.custom_from_mask

torch.nn.utils.prune.custom_from_mask(module, name, mask) [исходный код]

Обрезает тензор, соответствующий параметру с именем name в module, применяя предварительно вычисленную маску из mask.

Изменяет модуль на месте (и также возвращает изменённый модуль), выполняя следующие действия:

  1. добавляет именованный буфер с именем name+'_mask', соответствующий двоичной маске, применённой методом обрезки к параметру name.
  2. заменяет параметр name его обрезанной версией, сохраняя исходный (необрезанный) параметр в новом параметре с именем name+'_orig'.
Параметры:
  • module (nn.Module) – модуль, содержащий тензор для обрезки
  • name (str) – имя параметра в module, к которому будет применена обрезка.
  • mask (Tensor) – двоичная маска, применяемая к параметру.
Возвращает:

изменённую (то есть обрезанную) версию входного модуля

Тип возвращаемого значения:

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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API