Quickstart¶
Every snippet on this page runs as-is against the current version.
Important
Work in float64. All layers and models default to torch.float64, because
eigendecompositions lose accuracy quickly in float32. random_SPD defaults to
float32, so pass dtype=torch.float64 explicitly (mixing both raises a dtype
error in the first layer).
SPD matrices and matrix functions¶
import torch
from yetanotherspdnet.random.spd import random_SPD
from yetanotherspdnet.functions.spd_linalg import expm_symmetric, logm_SPD
generator = torch.Generator().manual_seed(0)
X = random_SPD(
n_features=8, n_matrices=32, cond=100, dtype=torch.float64, generator=generator
)
print(X.shape) # torch.Size([32, 8, 8]): a batch of 32 matrices 8 x 8
log_X = logm_SPD(X)[0] # matrix functions return (result, eigvals, eigvecs)
print(torch.allclose(expm_symmetric(log_X)[0], X)) # True
Every function accepts any number of leading batch dimensions
(..., n, n): (B, n, n) for a batch, (B, T, n, n) for sequences.
Riemannian means¶
The mean of SPD matrices depends on the geometry. Each geometry lives in its
own module under functions.spd_geometries:
from yetanotherspdnet.functions.spd_geometries.affine_invariant import (
affine_invariant_mean,
)
from yetanotherspdnet.functions.spd_geometries.bures_wasserstein import (
bures_wasserstein_mean,
)
from yetanotherspdnet.functions.spd_geometries.log_euclidean import log_euclidean_mean
G_ai = affine_invariant_mean(X, n_iterations=10) # Karcher flow
G_le = log_euclidean_mean(X) # closed form
G_bw = bures_wasserstein_mean(X, n_iterations=10) # fixed point
print(G_ai.shape) # torch.Size([8, 8])
See Geometries on the SPD manifold for the formulas and when to use which.
Layers¶
Layers are regular torch.nn.Modules and compose with torch.nn.Sequential:
from yetanotherspdnet.nn import BiMap, LogEig, ReEig, Vech
features = torch.nn.Sequential(
BiMap(8, 4), # 8x8 -> 4x4, orthonormal weight
ReEig(eps=1e-4), # clamp eigenvalues below eps
LogEig(), # SPD -> symmetric
Vech(), # 4x4 symmetric -> 10 coefficients
)
print(features(X).shape) # torch.Size([32, 10])
A complete model¶
from yetanotherspdnet import SPDnet
model = SPDnet(
input_dim=8,
hidden_layers_size=[6, 4], # BiMap 8->6->4
output_dim=3, # number of classes
batchnorm=True,
batchnorm_mean_type="affine_invariant",
)
print(model(X).shape) # torch.Size([32, 3])
It trains like any PyTorch model:
y = torch.randint(0, 3, (32,), generator=generator)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-2)
loss_fn = torch.nn.CrossEntropyLoss()
model.train()
for epoch in range(5):
optimizer.zero_grad()
loss = loss_fn(model(X), y)
loss.backward()
optimizer.step()
model.eval() # batch normalization now uses its running statistics
with torch.no_grad():
predictions = model(X).argmax(dim=-1)
The orthonormality of the BiMap weights and the positivity of the batch
normalization biases are handled by torch.nn.utils.parametrize, so a plain
Euclidean optimizer such as Adam or SGD is enough.
Batch normalization on its own¶
from yetanotherspdnet.nn import BatchNormSPDMeanScalarVariance
bn = BatchNormSPDMeanScalarVariance(n_features=8, mean_type="log_euclidean")
print(bn(X).shape) # torch.Size([32, 8, 8])
See Batch normalization for the available geometries and options.
Residual networks¶
from yetanotherspdnet import GBWBNRResNet, RResNet
rresnet = RResNet(
input_dim=8, hidden_layers_size=[6, 4], n_residual_blocks=[1, 1], output_dim=3
)
gbwbn = GBWBNRResNet(input_dim=8, hidden_dim=4, output_dim=3)
print(rresnet(X).shape, gbwbn(X).shape) # torch.Size([32, 3]) torch.Size([32, 3])