Group Regularize Callback

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.


source

GroupRegularizeCallback

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:

from fasterai.core.criteria import large_final
from fasterai.core.schedule import lin
from fasterai.prune.pruner import Pruner
from fasterai.regularize.group_regularize_callback import GroupRegularizeCallback

xb, _ = learn.dls.one_batch()
pruner = Pruner(learn.model, 0.3, 'local', large_final, reg=1e-4, example_inputs=xb)
learn.fit(10, cbs=GroupRegularizeCallback(pruner, schedule=lin))
pruner.prune_model()
learn.fit(2, reset_opt=True)
  • 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
  • Schedules - Ramp the penalty over training

Tests live in nbs/tests/test_group_regularize_callback.ipynb.