latest (dev)
Copy
Latest development documentation · Updated 2026-10-08
PruningContainer
- class tensorplay.ao.pruning.PruningContainer(*args: BasePruningMethod)[source]
Sequence of pruning methods applied iteratively to the same tensor.
Tracks the order in which methods were added and combines successive pruning calls: each new method only ranks the entries or channels that the previous masks left unpruned.
Accepts as argument an instance of a
BasePruningMethodor an iterable of them.- add_pruning_method(method: BasePruningMethod | None) None[source]
Add a child pruning
methodto the container.- Parameters:
method – child pruning method to be added to the container.
- Raises:
TypeError – if
methodis neitherNonenor aBasePruningMethodinstance.ValueError – if
methodacts on a different tensor name than the methods already held by the container.
- classmethod apply(module: Module, name: str, *args: Any, importance_scores: Tensor | None = None, **kwargs: Any) BasePruningMethod
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
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.
- compute_mask(t: Tensor, default_mask: Tensor) Tensor[source]
Apply the latest method and merge its mask into
default_mask.The new partial mask is computed on the entries or channels that
default_maskhas not zeroed out. Which portion oftthe new mask is derived from depends on thePRUNING_TYPEof the last method:'unstructured': the mask is computed from the flattened list of entries not yet masked;'structured': the mask is computed from the channels that still hold at least one unmasked entry;'global': the mask is computed across all entries.
- Parameters:
t – tensor representing the parameter to prune (same shape as
default_mask).default_mask – mask accumulated from previous pruning iterations.
- Returns:
The mask combining the effects of
default_maskand of the latest method, with the same shape asdefault_maskandt.
- prune(t: Tensor, default_mask: Tensor | None = None, importance_scores: Tensor | None = None) Tensor
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
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?

