Train with a Pruner’s group penalty before pruning
Overview
GroupRegularizeCallback calls Pruner.regularize() at every optimizer step of a fastai fit: the gradients receive the group penalty of a Pruner built with reg > 0, so training pushes the coupled channels the pruner will remove (a convolution’s filters, its BatchNorm scale, the input slices of the layers that read it) towards zero. The same Pruner then removes them with prune_model(), once the fit is over.
def GroupRegularizeCallback( pruner:Pruner, # Built with `reg > 0` on the model the Learner trains schedule:Schedule |None=None, # Ramps the penalty: `pruner.reg * progress` verbose:bool=False, # Print the penalty after each epoch):
Add pruner’s group penalty to the gradients at every optimizer step, so training pushes coupled channels towards zero before pruner.prune_model() removes them
Usage Example
Build the Pruner on learn.model with reg > 0, train with the callback, prune, then fine-tune:
The prune replaces the parameters of every layer it shrinks, so the next fit needs reset_opt=True: without it, fastai keeps the optimizer it built on the old parameters.
It works under learn.to_fp16() and GradientAccumulation: the penalty is multiplied by the loss scale before MixedPrecision unscales the gradients, and added once per optimizer step, not once per batch. NonNativeMixedPrecision is refused.
Frozen layers get no penalty.
Prune after the fit, not during: PruneCallback, which prunes during training, is refused in the same fit.
See Also
Pruner - reg, alpha and regularize(), the penalty this callback applies
PruneCallback - Prune during training instead of after it
RegularizeCallback - Penalize each layer’s own weights, one layer at a time