How to Add Batch Normalization to an SPDNet#

Insert Riemannian batch normalization into an existing SPDNet pipeline to stabilize training and improve convergence.

Prerequisites: Familiarity with SPDNet building blocks (see Building Blocks of SPD Neural Networks).

The Problem#

You have a working SPDNet but training is unstable or converges slowly. Adding batch normalization after each BiMap layer can help.

import torch
import torch.nn as nn

from spd_learn.modules import BiMap, LogEig, ReEig, SPDBatchNormLie

Step 1: Choose Your Normalization Layer#

spd_learn provides three batch normalization modules:

Module

When to Use

SPDBatchNormMeanVar

Standard choice. AIRM Frechet mean + variance scaling.

SPDBatchNormLie

Multiple metrics (AIM, LEM, LCM). Based on Lie group structure [Chen et al., 2024].

SPDBatchNormMean

Mean-only centering (no variance scaling). Simplest option.

Step 2: Insert After BiMap, Before ReEig#

The standard placement is BiMap -> BN -> ReEig. The final BiMap uses BN but skips ReEig:

dims = [64, 32, 16]  # your SPD matrix dimensions
layers = []
for i in range(len(dims) - 1):
    layers.append(BiMap(dims[i], dims[i + 1]))
    layers.append(SPDBatchNormLie(dims[i + 1], metric="LEM"))
    if i < len(dims) - 2:  # no ReEig after last BiMap
        layers.append(ReEig())

features = nn.Sequential(*layers)
print(features)
Sequential(
  (0): ParametrizedBiMap(
    (parametrizations): ModuleDict(
      (weight): ParametrizationList(
        (0): _Orthogonal()
      )
    )
  )
  (1): ParametrizedSPDBatchNormLie(
    num_features=32, metric=LEM, theta=1.0, alpha=1.0, beta=0.0, momentum=0.1, congruence=cholesky
    (parametrizations): ModuleDict(
      (bias): ParametrizationList(
        (0): SymmetricPositiveDefinite()
      )
      (shift): ParametrizationList(
        (0): PositiveDefiniteScalar()
      )
    )
  )
  (2): ReEig()
  (3): ParametrizedBiMap(
    (parametrizations): ModuleDict(
      (weight): ParametrizationList(
        (0): _Orthogonal()
      )
    )
  )
  (4): ParametrizedSPDBatchNormLie(
    num_features=16, metric=LEM, theta=1.0, alpha=1.0, beta=0.0, momentum=0.1, congruence=cholesky
    (parametrizations): ModuleDict(
      (bias): ParametrizationList(
        (0): SymmetricPositiveDefinite()
      )
      (shift): ParametrizationList(
        (0): PositiveDefiniteScalar()
      )
    )
  )
)

Step 3: Complete Network#

Wrap the features with LogEig and a linear classifier:

class SPDNetWithBN(nn.Module):
    """SPDNet with configurable batch normalization."""

    def __init__(self, dims, n_classes, metric="LEM"):
        super().__init__()
        layers = []
        for i in range(len(dims) - 1):
            layers.append(BiMap(dims[i], dims[i + 1]))
            layers.append(SPDBatchNormLie(dims[i + 1], metric=metric))
            if i < len(dims) - 2:
                layers.append(ReEig())
        self.features = nn.Sequential(*layers)
        self.logeig = LogEig(upper=False, flatten=True)
        self.classifier = nn.Linear(dims[-1] ** 2, n_classes)

    def forward(self, x):
        return self.classifier(self.logeig(self.features(x)))


model = SPDNetWithBN([64, 32, 16], n_classes=4, metric="LEM")

Verify it works with a dummy forward pass:

X = torch.randn(8, 64, 64)
X = X @ X.mT + 0.01 * torch.eye(64)
out = model(X)
print(f"Input: {X.shape} -> Output: {out.shape}")  # [8, 64, 64] -> [8, 4]
Input: torch.Size([8, 64, 64]) -> Output: torch.Size([8, 4])

Key Points#

  • Place BN after BiMap and before ReEig

  • Use model.train() / model.eval() – BN uses running stats at inference time

  • momentum=0.1 (default) works well in most cases

  • Consider float64 for numerical stability with Riemannian operations

See also

References#

[1]

Ziheng Chen, Yue Song, Yunmei Xu, and Nicu Sebe. A lie group approach to riemannian batch normalization. In International Conference on Learning Representations. 2024. URL: https://openreview.net/forum?id=okYdj8Ysru.

Total running time of the script: (0 minutes 0.016 seconds)