latest (dev)
Copy
Latest development documentation · Updated 2026-10-08
BasePruningMethod
- class tensorplay.ao.pruning.BasePruningMethod[source]
Abstract base class for pruning techniques.
Subclasses must override
compute_mask()and declare aPRUNING_TYPEclass attribute (one of'unstructured','structured'or'global'). The classmethodapply()installs the reparameterization on a module and registers an instance of the subclass as a forward pre-hook.- classmethod apply(module: Module, name: str, *args: Any, importance_scores: Tensor | None = None, **kwargs: Any) BasePruningMethod[source]
Install the pruning reparameterization for
module[name].Moves the original parameter to
name + '_orig', registers the computed mask as the buffername + '_mask', stores the masked values undernameand adds a forward pre-hook that re-applies the mask on every forward pass. If the tensor is already pruned by another method, the new method is composed with the existing one through aPruningContainer.- Parameters:
module – module containing the tensor to prune.
name – parameter name within
moduleon which pruning acts.args – positional arguments forwarded to the subclass constructor.
importance_scores – tensor of importance scores with the same shape as
module[name]; each entry ranks the corresponding element of the parameter. When unspecified, the parameter itself is used as its own importance scores.kwargs – keyword arguments forwarded to the subclass constructor.
- Returns:
The pruning method (or container of methods) now attached to the module.
- apply_mask(module: Module) Tensor[source]
Return the pruned version of the tensor held by
module.Fetches the mask and the original values from the module and returns their elementwise product.
- Parameters:
module – module holding the pruned tensor.
- Returns:
The product of the mask and the original values.
- abstractmethod compute_mask(t: Tensor, default_mask: Tensor) Tensor[source]
Compute the pruning mask for the input tensor
t.Starting from
default_mask(a mask of ones whenthas never been pruned), derive the new mask according to the recipe of the concrete method. Entries already zeroed bydefault_maskmust stay zero.- Parameters:
t – tensor whose entries or channels are ranked for pruning, typically the importance scores of the parameter.
default_mask – mask accumulated from previous pruning iterations; must be respected by the new mask. Same shape as
t.
- Returns:
The mask to apply to
t, with the same shape ast.
- prune(t: Tensor, default_mask: Tensor | None = None, importance_scores: Tensor | None = None) Tensor[source]
Return a pruned copy of the input tensor
t.Applies the rule implemented by
compute_mask()without any module-side reparameterization.- Parameters:
t – tensor to prune (same shape as
default_mask).importance_scores – tensor of importance scores with the same shape as
t; each entry ranks the corresponding element oft. When unspecified,titself is used.default_mask – mask from a previous pruning iteration, if any. Pruning must respect the entries it already zeroes. When unspecified, a mask of ones is used.
- Returns:
The pruned version of
t.
- remove(module: Module) None[source]
Make the current pruning of
modulepermanent.The pruned values remain pruned: the product of mask and original values is written back into the parameter
name, and the auxiliary parametername + '_orig'and buffername + '_mask'are dropped.Note
Pruning itself is NOT undone or reversed!
Help improve this page
Found an error, an unclear step, or a missing example?

