Spec-Zone.ru › PyTorch 2

RandomStructured

class torch.nn.utils.prune.RandomStructured(amount, dim=-1) [source]

Случайное обрезание целых (в данный момент необрезанных) каналов в тензоре.

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

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

Параметры
  • 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) [source]

Вычисляет и возвращает маску для входного тензора 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' удаляется из буферов.

Примечание

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

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

Spec-Zone.ru

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