Spec-Zone.ru › PyTorch 2

СлучайныйНеструктурированный

class torch.nn.utils.prune.RandomUnstructured(amount) [source]

Случайное обрезание (в настоящее время не обрезанных) элементов тензора.

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

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

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

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

Параметры

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

Возвращает

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

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

pruned_tensor (torch.Tensor)

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.RandomUnstructured.html

Spec-Zone.ru

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