Spec-Zone.ru › PyTorch 2.14

RandomStructured

class torch.nn.utils.prune.RandomStructured(amount, dim=-1) [исходный код]

Случайным образом исключает из тензора целые каналы (которые ещё не были исключены).

Параметры:
  • amount (int или float) – количество параметров, которые нужно исключить. Если float, должно быть в диапазоне от 0.0 до 1.0 и представлять долю параметров для исключения. Если int, представляет абсолютное количество параметров для исключения.
  • dim (int, необязательно) – индекс размерности, вдоль которой определяются исключаемые каналы. По умолчанию: -1.
classmethod apply(module, name, amount, dim=-1) [исходный код]

Добавляет обрезку на лету и переопределение параметризации тензора.

Добавляет предварительный перехватчик прямого прохода, который включает обрезку на лету и переопределение параметризации тензора через исходный тензор и маску обрезки.

Параметры:
  • module (nn.Module) – модуль, содержащий тензор для обрезки
  • name (str) – имя параметра в module, к которому будет применяться обрезка.
  • amount (int или float) – количество параметров, которые нужно исключить. Если float, должно быть в диапазоне от 0.0 до 1.0 и представлять долю параметров для исключения. Если int, представляет абсолютное количество параметров для исключения.
  • dim (int, необязательно) – индекс размерности, вдоль которой определяются исключаемые каналы. По умолчанию: -1.
apply_mask(module) [исходный код]

Просто выполняет умножение обрезаемого параметра на созданную маску.

Извлекает маску и исходный тензор из модуля и возвращает обрезанную версию тензора.

Параметры:

module (nn.Module) – модуль, содержащий тензор для обрезки

Возвращает:

обрезанную версию входного тензора

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

pruned_tensor (torch.Tensor)

compute_mask(t, default_mask) [исходный код]

Вычисляет и возвращает маску для входного тензора t.

На основе исходной default_mask (которая должна быть маской из единиц, если тензор ещё не обрезался) создаёт случайную маску, накладываемую на default_mask, случайным образом обнуляя каналы вдоль указанной размерности тензора.

Параметры:
  • t (torch.Tensor) – тензор, представляющий обрезаемый параметр
  • default_mask (torch.Tensor) – исходная маска предыдущих итераций обрезки, которую необходимо учитывать при применении новой маски. Имеет те же размерности, что и t.
Возвращает:

маску для применения к t с теми же размерностями, что и t

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

mask (torch.Tensor)

Вызывает исключение:

IndexError – если self.dim >= len(t.shape)

prune(t, default_mask=None, importance_scores=None) [исходный код]

Вычисляет и возвращает обрезанную версию входного тензора t.

В соответствии с правилом обрезки, указанным в compute_mask().

Параметры:
  • t (torch.Tensor) – тензор для обрезки (с теми же размерностями, что и default_mask).
  • importance_scores (torch.Tensor) – тензор с оценками важности (той же формы, что и t), используемый для вычисления маски обрезки t. Значения в этом тензоре указывают важность соответствующих элементов в t, который обрезается. Если не указан или равен None, вместо него будет использоваться тензор t.
  • default_mask (torch.Tensor, необязательно) – маска предыдущей итерации обрезки, если она есть. Учитывается при определении части тензора, к которой следует применить обрезку. Если None, по умолчанию используется маска из единиц.
Возвращает:

обрезанную версию тензора t.

remove(module) [исходный код]

Удаляет переопределение параметризации обрезки из модуля.

Обрезанный параметр с именем name остаётся обрезанным навсегда, а параметр с именем name+'_orig' удаляется из списка параметров. Аналогично, буфер с именем name+'_mask' удаляется из списка буферов.

Примечание

Сама обрезка НЕ отменяется и НЕ обращается!

© 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.RandomStructured.html

Spec-Zone.ru

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