TensorPlay
latest (dev)
Copy
View Markdown

Latest development documentation · Updated 2026-10-08

tensorplay.ao.pruning.global_unstructured

tensorplay.ao.pruning.global_unstructured(parameters: Iterable[tuple[Module, str]], pruning_method: type[BasePruningMethod], importance_scores: dict[tuple[Module, str], Tensor] | None = None, **kwargs: Any) → None[source]

Prune several tensors jointly under a single unstructured budget.

Aggregates the importance scores of every listed parameter into one vector, computes a single mask under the shared amount budget, and slices that mask back onto each parameter. Modifies the modules in place by:

  1. adding a named buffer called name + '_mask' for every listed parameter;

  2. replacing each parameter name by its masked version, while the original (unmasked) values are stored in a new parameter named name + '_orig'.

Parameters:
  • parameters – iterable of (module, name) tuples identifying the parameters to prune globally, i.e. by aggregating all values before deciding which units to remove.

  • pruning_method – a pruning method class from this package (or a user-defined subclass of BasePruningMethod) whose PRUNING_TYPE is 'unstructured'.

  • importance_scores – mapping from (module, name) tuples to the corresponding importance-scores tensor (same shape as the parameter). Parameters absent from the mapping use their own values as importance scores.

  • kwargs – keyword arguments forwarded to pruning_method, typically amount: the quantity of units to prune across all listed parameters (a float fraction in [0, 1] or an absolute int).

Raises:

TypeError – if parameters is not an iterable, if importance_scores is not a dict, or if the PRUNING_TYPE of pruning_method is not 'unstructured'.

Note

Global pruning is restricted to unstructured methods: a structured norm is only comparable across channels of equal size, which cannot be guaranteed across heterogeneous parameters.

Examples

>>> # xdoctest: +SKIP
>>> net = nn.Sequential(nn.Linear(10, 4), nn.Linear(4, 1))
>>> parameters_to_prune = (
...     (net[0], "weight"),
...     (net[1], "weight"),
... )
>>> global_unstructured(
...     parameters_to_prune,
...     pruning_method=L1Unstructured,
...     amount=10,
... )

On this page

Ask DeepWiki