RandomStructured
-
class torch.nn.utils.prune.RandomStructured(amount, dim=-1)[исходный код] -
Случайным образом исключает из тензора целые каналы (которые ещё не были исключены).
- Параметры:
-
-
amount (int или float) – количество параметров, которые нужно исключить. Если
float, должно быть в диапазоне от 0.0 до 1.0 и представлять долю параметров для исключения. Еслиint, представляет абсолютное количество параметров для исключения. - dim (int, необязательно) – индекс размерности, вдоль которой определяются исключаемые каналы. По умолчанию: -1.
-
amount (int или float) – количество параметров, которые нужно исключить. Если
-
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 (torch.Tensor) – тензор для обрезки (с теми же размерностями, что и
- Возвращает:
-
обрезанную версию тензора
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