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.
Riemannian batch normalization of SPD matrices: mean centring and bias. |
|
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,
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 featuresmean_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 Nonemomentum (
float, optional) – Momentum for running mean update. Default is 0.01norm_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.01minibatch_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 100parametrization (
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 100use_autograd (
bool, optional) – Use torch autograd for the computation of the gradient rather than the analytical formula. Default is Falsedevice (
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.Tensorofshape (n_features,n_features)) – Learnable SPD bias \(G\) (parametrized; trainable tensorparametrizations.Covbias.original).running_mean (
torch.Tensorofshape (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 ofrunning_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.Tensorofshape (...,n_features,n_features)) – Batch of SPD matrices- Returns:
data_transformed (
torch.Tensorofshape (...,n_features,n_features)) – Batch of transformed (normalized then biased) SPD matrices- Return type:
- 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,
Riemannian batch normalization of SPD matrices: mean, scalar dispersion, bias.
Extends
BatchNormSPDMeanwith 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. Withmean_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 featuresmean_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 Nonemomentum (
float, optional) – Momentum for running mean update. Default is 0.01norm_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.01minibatch_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 100parametrization (
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 100use_autograd (
bool, optional) – Use torch autograd for the computation of the gradient rather than the analytical formula. Default is Falsebw_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.5bw_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 Truedevice (
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.Tensorofshape (n_features,n_features)) – Learnable SPD bias.stdScalarbias (
torch.Tensor) – Learnable positive scalar scale \(s\).running_mean (
torch.Tensorofshape (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.Tensorofshape (...,n_features,n_features)) – Batch of SPD matrices- Returns:
data_transformed (
torch.Tensorofshape (...,n_features,n_features)) – Batch of transformed (normalized then biased) SPD matrices- Return type: