Note
Go to the end to download the full example code.
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 |
|---|---|
Standard choice. AIRM Frechet mean + variance scaling. |
|
Multiple metrics (AIM, LEM, LCM). Based on Lie group structure [Chen et al., 2024]. |
|
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:
Input: torch.Size([8, 64, 64]) -> Output: torch.Size([8, 4])
Key Points#
Place BN after
BiMapand beforeReEigUse
model.train()/model.eval()– BN uses running stats at inference timemomentum=0.1(default) works well in most casesConsider
float64for numerical stability with Riemannian operations
See also
Batch Normalization on SPD Manifolds – Learn how BN works on SPD manifolds
How to Choose a Metric for Batch Normalization – Choosing between AIM, LEM, and LCM
SPDBatchNormLie– API referenceSPDBatchNormMeanVar– API reference
References#
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)