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.
Parametrization mapping a symmetric matrix to an SPD matrix. |
|
SPD parametrization around a moving reference point. |
|
Stiefel parametrization around a moving reference point. |
|
Parametrization constraining a scalar to be positive with a softplus. |
|
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:
- forward(tangent_vector)[source]¶
Mapping from the tangent space at identity to the SPD manifold
- Parameters:
tangent_vector (
torch.Tensorofshape (n_features,n_features)) – Symmetric matrix- Returns:
spd_matrix (
torch.Tensorofshape (n_features,n_features)) – SPD matrix- Return type:
- right_inverse(spd_matrix)[source]¶
Mapping from the SPD manifold to the tangent space at identity
- Parameters:
spd_matrix (
torch.Tensorofshape (n_features,n_features)) – SPD matrix- Returns:
tangent_vector (
torch.Tensorofshape (n_features,n_features)) – Symmetric matrix- Return type:
- class SPDAdaptiveParametrization(
- n_features,
- initial_reference=None,
- mapping='softplus',
- use_autograd=False,
- device=device(type='cpu'),
- dtype=torch.float64,
SPD parametrization around a moving reference point.
The trainable tensor is a tangent vector at the current reference point. Calling
update_reference_pointmoves 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 featuresinitial_reference (
torch.Tensor | None, optional) – Initial reference point. If None, the identity matrix is selected. Default is Nonemapping (
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 Falsedevice (
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.Tensorofshape (n_features,n_features)) – Symmetric matrix- Returns:
spd_matrix (
torch.Tensorofshape (n_features,n_features))- Return type:
- right_inverse(spd_matrix)[source]¶
Mapping from SPD manifold onto the tangent space at reference_point
- Parameters:
spd_matrix (
torch.Tensorofshape (n_features,n_features)) – SPD matrix- Returns:
tangent_vector (
torch.Tensorofshape (n_features,n_features)) – Symmetric matrix- Return type:
- class StiefelAdaptiveParametrization(
- n_in,
- n_out,
- initial_reference=None,
- mapping='QR',
- use_autograd=False,
- device=device(type='cpu'),
- dtype=torch.float64,
- generator=None,
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 rowsn_out (
int) – Number of columnsinitial_reference (
torch.Tensor | None, optional) – Initial reference point. If None, a random point on Stiefel is generated. Default is Nonemapping (
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 Falseuse_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 Falsedevice (
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.float64generator (
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.Tensorofshape (n_in,n_out)) – Rectangular matrix, tangent vector at reference_point- Returns:
weight (
torch.Tensorofshape (n_in,n_out)) – Orthogonal matrix- Return type:
- right_inverse(weight)[source]¶
Mapping from Stiefel manifold to the tangent space at reference_point (achieved through orthogonal projection)
- Parameters:
weight (
torch.Tensorofshape (n_in,n_out)) – Orthogonal matrix- Returns:
weight_tangent (
torch.Tensorofshape (n_in,n_out)) – Rectangular matrix, tangent vector at reference_point- Return type:
- 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.Tensorofshape ()) – Real scalar- Returns:
scalar_pd (
torch.Tensorofshape ()) – Positive definite scalar- Return type:
- right_inverse(scalar_pd)[source]¶
Mapping from positive definite scalar onto real scalars through the inverse SoftPlus function
- Parameters:
scalar_pd (
torch.Tensorofshape ()) – Positive definite scalar- Returns:
scalar (
torch.Tensorofshape ()) – Real scalar- Return type:
- 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.Tensorofshape ()) – Real scalar (unconstrained)- Returns:
scalar_constrained (
torch.Tensorofshape ()) – Scalar in [0, 1]- Return type:
- right_inverse(scalar_constrained)[source]¶
Mapping from [0, 1] to real scalars through inverse sigmoid (logit)
- Parameters:
scalar_constrained (
torch.Tensorofshape ()) – Scalar in [0, 1]- Returns:
scalar (
torch.Tensorofshape ()) – Real scalar (unconstrained)- Return type: