Batch normalization

yetanotherspdnet.nn.batchnorm — centring (and optionally rescaling) a batch of SPD matrices around a learned bias, in any of the Geometries. The Batch normalization guide describes the steps, the choice of mean, the training / evaluation behaviour and the GBWBN options.

BatchNormSPDMean

Riemannian batch normalization of SPD matrices: mean centring and bias.

BatchNormSPDMeanScalarVariance

Riemannian batch normalization of SPD matrices: mean, scalar dispersion, bias.

class BatchNormSPDMean(
n_features,
mean_type='affine_invariant',
mean_options=None,
momentum=0.01,
norm_strategy='classical',
minibatch_mode='constant',
minibatch_momentum=0.01,
minibatch_maxstep=100,
parametrization='softplus',
parametrization_mode='static',
n_steps_ref_update=100,
use_autograd=False,
device=device(type='cpu'),
dtype=torch.float64,
)[source]

Riemannian batch normalization of SPD matrices: mean centring and bias.

In training mode each batch is centred on its SPD mean \(\bar X\) (computed in the geometry selected by mean_type) and re-biased by a learnable SPD matrix \(G\); in evaluation mode the running mean replaces the batch mean. A batch holding a single matrix has no batch statistics and is normalized with the running mean, which is then left unchanged.

Batch normalization layer for SPDnet relying on a SPD mean. Only the SPD mean is normalized

Parameters:
  • n_features (int) – Number of features

  • mean_type (str, optional) – Choice of SPD mean. Default is “affine_invariant”. Choices are: “affine_invariant”, “log_euclidean”, “arithmetic”, “harmonic”, “geometric_arithmetic_harmonic”, “bures_wasserstein”

  • mean_options (dict | None, optional) – Options for the SPD mean computation. For affine-invariant mean, one can typically set {‘n_iterations’: 5}. Currently, for others, no options available. Default is None

  • momentum (float, optional) – Momentum for running mean update. Default is 0.01

  • norm_strategy (str, optional) – Strategy for normalization. Default is “classical”. Choices are: “classical” and “minibatch”

  • minibatch_mode (str, optional) – How the minibatch momentum behaves during the training. Default is “constant”. Choices are: “constant”, “decay”, “growth”

  • minibatch_momentum (float, optional) – Momentum for mean regularization in minibatch normalization strategy. If minibatch_mode is “decay”, this momentum corresponds to the minimum momentum that is reached. If minibatch_mode is “growth”, this momentum corresponds to the initial momentum. Default is 0.01

  • minibatch_maxstep (int, optional) – If minibatch_mode is “decay” or “growth”, this is the training step at which the minibatch momentum attains its final value. Default is 100

  • parametrization (str, optional) – Parametrization to apply on covariance bias. Default is “softplus”. Choices are: “softplus”, “exp”

  • parametrization_mode (str, optional) – Parametrization mode. Default is “static”. Choices are: “static” and “dynamic”

  • n_steps_ref_update (int, optional) – If parametrization_mode is “dynamic”, number of steps in between each reference point update. Default is 100

  • use_autograd (bool, optional) – Use torch autograd for the computation of the gradient rather than the analytical formula. Default is False

  • device (torch.device, optional) – Device on which to store the parameters. Default is torch.device(“cpu”)

  • dtype (torch.dtype, optional) – Data type of the layer. Default is torch.float64

Variables:
  • Covbias (torch.Tensor of shape (n_features, n_features)) – Learnable SPD bias \(G\) (parametrized; trainable tensor parametrizations.Covbias.original).

  • running_mean (torch.Tensor of shape (n_features, n_features)) – Buffer: running SPD mean used in evaluation mode.

  • mean_type (str) – Geometry of the mean: "affine_invariant", "log_euclidean", "arithmetic", "harmonic", "geometric_arithmetic_harmonic", "adaptive_geometric_arithmetic_harmonic" or "bures_wasserstein".

  • momentum (float) – Update rate of running_mean.

  • training_step (int) – Number of training batches seen (drives the "minibatch" strategy).

forward(data)[source]

Forward pass of the BatchNormSPDMean layer

Parameters:

data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices

Returns:

data_transformed (torch.Tensor of shape (..., n_features, n_features)) – Batch of transformed (normalized then biased) SPD matrices

Return type:

Tensor

class BatchNormSPDMeanScalarVariance(
n_features,
mean_type='affine_invariant',
mean_options=None,
momentum=0.01,
norm_strategy='classical',
minibatch_mode='constant',
minibatch_momentum=0.01,
minibatch_maxstep=100,
parametrization='softplus',
parametrization_mode='static',
n_steps_ref_update=100,
use_autograd=False,
bw_theta=0.5,
bw_batch_stats_grad=True,
device=device(type='cpu'),
dtype=torch.float64,
)[source]

Riemannian batch normalization of SPD matrices: mean, scalar dispersion, bias.

Extends BatchNormSPDMean with a scalar dispersion: centred matrices are rescaled along geodesics from the identity by \(s / \sigma\), where \(\sigma\) is the batch (or running) standard deviation and \(s\) a learnable positive scalar. With mean_type="bures_wasserstein" the layer implements GBWBN (generalized Bures-Wasserstein batch normalization) with learnable pre/post transforms.

Batch normalization layer for SPDnet relying on a SPD mean. Both the SPD mean and scalar variance are normalized

Parameters:
  • n_features (int) – Number of features

  • mean_type (str, optional) – Choice of SPD mean. Default is “affine_invariant”. Choices are: “affine_invariant”, “log_euclidean”, “arithmetic”, “harmonic”, “geometric_arithmetic_harmonic”, “bures_wasserstein”

  • mean_options (dict | None, optional) – Options for the SPD mean computation. For affine-invariant mean, one can typically set {‘n_iterations’: 5}. For bures_wasserstein, one can set {‘n_iterations’: 1}. Default is None

  • momentum (float, optional) – Momentum for running mean update. Default is 0.01

  • norm_strategy (str, optional) – Strategy for normalization. Default is “classical”. Choices are: “classical” and “minibatch”

  • minibatch_mode (str, optional) – How the minibatch momentum behaves during the training. Default is “constant”. Choices are: “constant”, “decay”, “growth”

  • minibatch_momentum (float, optional) – Momentum for mean regularization in minibatch normalization strategy. If minibatch_mode is “decay”, this momentum corresponds to the minimum momentum that is reached. If minibatch_mode is “growth”, this momentum corresponds to the initial momentum. Default is 0.01

  • minibatch_maxstep (int, optional) – If minibatch_mode is “decay” or “growth”, this is the training step at which the minibatch momentum attains its final value. Default is 100

  • parametrization (str, optional) – Parametrization to apply on covariance bias. Default is “softplus”. Choices are: “softplus”, “exp”

  • parametrization_mode (str, optional) – Parametrization mode. Default is “static”. Choices are: “static” and “dynamic”

  • n_steps_ref_update (int, optional) – If parametrization_mode is “dynamic”, number of steps in between each reference point update. Default is 100

  • use_autograd (bool, optional) – Use torch autograd for the computation of the gradient rather than the analytical formula. Default is False

  • bw_theta (float, optional) – Power deformation \(\theta\) of the generalized Bures-Wasserstein metric (GBWBN, Wang et al. 2025): data are mapped by \(X \mapsto M^{-1/2} X^\theta M^{-1/2}\) and the variance is measured with the deformed metric (divided by \(\theta^2\)). Only used when mean_type=”bures_wasserstein”. theta=1.0 is the plain BW metric; 0.5 is the best value of the paper’s ablation (HDM05: 69.1% vs 62.0% for theta=1) and the reference code’s default. Default is 0.5

  • bw_batch_stats_grad (bool, optional) – GBWBN only. If True (as in the paper), gradients flow through the batch mean and variance; if False (as in the reference code), they are computed without gradient, as constants. Default is True

  • device (torch.device, optional) – Device on which to store the parameters. Default is torch.device(“cpu”)

  • dtype (torch.dtype, optional) – Data type of the layer. Default is torch.float64

Variables:
  • Covbias (torch.Tensor of shape (n_features, n_features)) – Learnable SPD bias.

  • stdScalarbias (torch.Tensor) – Learnable positive scalar scale \(s\).

  • running_mean (torch.Tensor of shape (n_features, n_features)) – Buffer: running SPD mean.

  • running_std_scalar (torch.Tensor) – Buffer: running scalar standard deviation.

  • bw_G (bw_M,) – GBWBN only: learnable SPD pre-transform \(M\) and bias \(G\), applied in the transformed space as \(\hat G = M^{-1/2} G^\theta M^{-1/2}\).

  • bw_theta (float) – GBWBN only: power of the pre-transform \(X \mapsto X^\theta\).

forward(data)[source]

Forward pass of the BatchNormSPDMeanScalarVariance layer

Parameters:

data (torch.Tensor of shape (..., n_features, n_features)) – Batch of SPD matrices

Returns:

data_transformed (torch.Tensor of shape (..., n_features, n_features)) – Batch of transformed (normalized then biased) SPD matrices

Return type:

Tensor