Parametrizations

yetanotherspdnet.nn.parametrizations — modules registered with torch.nn.utils.parametrize to keep a parameter on its manifold while a standard optimizer updates an unconstrained tensor. The adaptive variants work around a reference point that moves during training (the "dynamic" mode of BiMap and of the batch normalization bias); see Layers and parametrizations.

SPDParametrization

Parametrization mapping a symmetric matrix to an SPD matrix.

SPDAdaptiveParametrization

SPD parametrization around a moving reference point.

StiefelAdaptiveParametrization

Stiefel parametrization around a moving reference point.

ScalarSoftPlusParametrization

Parametrization constraining a scalar to be positive with a softplus.

ScalarSigmoidParametrization

Parametrization to constrain a scalar to the interval [0, 1] using sigmoid.

class SPDParametrization(mapping='softplus', use_autograd=False)[source]

Parametrization mapping a symmetric matrix to an SPD matrix.

The eigenvalues of the unconstrained symmetric matrix go through a softplus (mapping="softplus") or an exponential (mapping="exp"). Used for the SPD biases of batch normalization.

SPD Parametrization

Parameters:
  • mapping (str, optional) – Mapping to obtain a SPD point from a symmetric matrix. Default is “softplus”. Choices are: “softplus” and “exp”

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

forward(tangent_vector)[source]

Mapping from the tangent space at identity to the SPD manifold

Parameters:

tangent_vector (torch.Tensor of shape (n_features, n_features)) – Symmetric matrix

Returns:

spd_matrix (torch.Tensor of shape (n_features, n_features)) – SPD matrix

Return type:

Tensor

right_inverse(spd_matrix)[source]

Mapping from the SPD manifold to the tangent space at identity

Parameters:

spd_matrix (torch.Tensor of shape (n_features, n_features)) – SPD matrix

Returns:

tangent_vector (torch.Tensor of shape (n_features, n_features)) – Symmetric matrix

Return type:

Tensor

class SPDAdaptiveParametrization(
n_features,
initial_reference=None,
mapping='softplus',
use_autograd=False,
device=device(type='cpu'),
dtype=torch.float64,
)[source]

SPD parametrization around a moving reference point.

The trainable tensor is a tangent vector at the current reference point. Calling update_reference_point moves the reference to the current value, which keeps the chart well conditioned during long trainings (parametrization_mode="dynamic" in the layers).

Adaptive SPD Parametrization

Parameters:
  • n_features (int) – Number of features

  • initial_reference (torch.Tensor | None, optional) – Initial reference point. If None, the identity matrix is selected. Default is None

  • mapping (str, optional) – Mapping to obtain a SPD point from a tangent vector. Default is “softplus”. Choices are: “softplus” and “exp”

  • use_autograd (bool | dict, optional) – Use torch autograd for gradient computation. Can be bool for all layers, or dict with keys: ‘bimap’, ‘reeig’, ‘logeig’, ‘batchnorm’, ‘vec’. Note that Vech module always uses manual gradient. Default is False

  • device (torch.device, optional) – Device to run model on. Default is torch.device(‘cpu’)

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

forward(tangent_vector)[source]

Mapping from the tangent space at reference_point onto the SPD manifold

Parameters:

tangent_vector (torch.Tensor of shape (n_features, n_features)) – Symmetric matrix

Returns:

spd_matrix (torch.Tensor of shape (n_features, n_features))

Return type:

Tensor

right_inverse(spd_matrix)[source]

Mapping from SPD manifold onto the tangent space at reference_point

Parameters:

spd_matrix (torch.Tensor of shape (n_features, n_features)) – SPD matrix

Returns:

tangent_vector (torch.Tensor of shape (n_features, n_features)) – Symmetric matrix

Return type:

Tensor

update_reference_point()[source]

Update reference point with last SPD value

class StiefelAdaptiveParametrization(
n_in,
n_out,
initial_reference=None,
mapping='QR',
use_autograd=False,
device=device(type='cpu'),
dtype=torch.float64,
generator=None,
)[source]

Stiefel parametrization around a moving reference point.

The trainable tensor is a tangent vector at the current reference point, mapped back to the Stiefel manifold by a QR or polar retraction. Used by BiMap with parametrization_mode="dynamic".

Adaptive Stiefel Parametrization

Parameters:
  • n_in (int) – Number of rows

  • n_out (int) – Number of columns

  • initial_reference (torch.Tensor | None, optional) – Initial reference point. If None, a random point on Stiefel is generated. Default is None

  • mapping (str, optional) – Mapping to obtain a point on Stiefel from a tangent vector. Default is “QR”. Choices are: “QR” and “polar” WARNING: with “polar”, use_autograd needs to be False

  • use_autograd (bool | dict, optional) – Use torch autograd for gradient computation. Can be bool for all layers, or dict with keys: ‘bimap’, ‘reeig’, ‘logeig’, ‘batchnorm’, ‘vec’. Note that Vech module always uses manual gradient. Default is False

  • device (torch.device, optional) – Device to run model on. Default is torch.device(‘cpu’)

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

  • generator (torch.Generator, optional) – Generator to ensure reproducibility. Default is None

forward(weight_tangent)[source]

Mapping from the tangent space of reference_point to the Stiefel manifold

Parameters:

weight_tangent (torch.Tensor of shape (n_in, n_out)) – Rectangular matrix, tangent vector at reference_point

Returns:

weight (torch.Tensor of shape (n_in, n_out)) – Orthogonal matrix

Return type:

Tensor

right_inverse(weight)[source]

Mapping from Stiefel manifold to the tangent space at reference_point (achieved through orthogonal projection)

Parameters:

weight (torch.Tensor of shape (n_in, n_out)) – Orthogonal matrix

Returns:

weight_tangent (torch.Tensor of shape (n_in, n_out)) – Rectangular matrix, tangent vector at reference_point

Return type:

Tensor

update_reference_point()[source]

Update reference point with last Stiefel value

class ScalarSoftPlusParametrization(*args, **kwargs)[source]

Parametrization constraining a scalar to be positive with a softplus.

Initialize internal Module state, shared by both nn.Module and ScriptModule.

forward(scalar)[source]

Positive definite scalars parametrization using the SoftPlus function (rescaled so that f(0) = 1 as compared to default torch function)

Parameters:

scalar (torch.Tensor of shape ()) – Real scalar

Returns:

scalar_pd (torch.Tensor of shape ()) – Positive definite scalar

Return type:

Tensor

right_inverse(scalar_pd)[source]

Mapping from positive definite scalar onto real scalars through the inverse SoftPlus function

Parameters:

scalar_pd (torch.Tensor of shape ()) – Positive definite scalar

Returns:

scalar (torch.Tensor of shape ()) – Real scalar

Return type:

Tensor

class ScalarSigmoidParametrization(*args, **kwargs)[source]

Parametrization to constrain a scalar to the interval [0, 1] using sigmoid.

Initialize internal Module state, shared by both nn.Module and ScriptModule.

forward(scalar)[source]

Mapping from real scalar to [0, 1] through sigmoid function

Parameters:

scalar (torch.Tensor of shape ()) – Real scalar (unconstrained)

Returns:

scalar_constrained (torch.Tensor of shape ()) – Scalar in [0, 1]

Return type:

Tensor

right_inverse(scalar_constrained)[source]

Mapping from [0, 1] to real scalars through inverse sigmoid (logit)

Parameters:

scalar_constrained (torch.Tensor of shape ()) – Scalar in [0, 1]

Returns:

scalar (torch.Tensor of shape ()) – Real scalar (unconstrained)

Return type:

Tensor