"""Preconditioner for pinns."""
from typing import Callable, cast
import torch
from scimba_torch.approximation_space.abstract_space import AbstractApproxSpace
from scimba_torch.numerical_solvers.functional_operator import (
ACCEPTED_PDE_TYPES,
TYPE_DICT_OF_VMAPS,
FunctionalOperator,
)
from scimba_torch.numerical_solvers.preconditioner_pinns import MatrixPreconditionerPinn
def _solve_cholesky(gram: torch.Tensor, grad: torch.Tensor) -> torch.Tensor:
"""Solve ``gram @ x = grad`` through Cholesky factorization (gram is SPD).
The ENG Gram matrix is SPD once regularized, so this is the cheapest exact
solver: O(m^3/3) flops vs O(4m^3/3) for a QR-based least-squares solve.
Falls back to an LU solve if the factorization fails (``gram`` not
numerically positive-definite -- can happen with adaptive regularization
pushed close to its floor), so switching the default solver away from the
always-robust ``"lstsq"`` doesn't introduce a crash risk.
Args:
gram: The (regularized) Gram matrix, a symmetric positive-definite
``(ndof, ndof)`` matrix.
grad: The right-hand side, a ``(ndof,)`` or ``(ndof, k)`` tensor.
Returns:
The solution ``x``, with the same shape as ``grad``.
"""
chol, info = torch.linalg.cholesky_ex(gram)
if bool(torch.any(info != 0)):
return _solve_lu(gram, grad)
rhs = grad if grad.ndim > 1 else grad.unsqueeze(-1)
sol = torch.cholesky_solve(rhs, chol)
return sol if grad.ndim > 1 else sol.squeeze(-1)
def _solve_lu(gram: torch.Tensor, grad: torch.Tensor) -> torch.Tensor:
"""Solve ``gram @ x = grad`` through LU factorization with partial pivoting.
Args:
gram: The (regularized) Gram matrix, a ``(ndof, ndof)`` matrix.
grad: The right-hand side, a ``(ndof,)`` or ``(ndof, k)`` tensor.
Returns:
The solution ``x``, with the same shape as ``grad``.
"""
return torch.linalg.solve(gram, grad)
def _solve_lstsq(gram: torch.Tensor, grad: torch.Tensor) -> torch.Tensor:
"""Solve ``gram @ x = grad`` in the least-squares sense (historical behavior).
Args:
gram: The (regularized) Gram matrix, a ``(ndof, ndof)`` matrix.
grad: The right-hand side, a ``(ndof,)`` or ``(ndof, k)`` tensor.
Returns:
The solution ``x``, with the same shape as ``grad``.
"""
return torch.linalg.lstsq(gram, grad).solution
# Registry of dense solvers for the natural-gradient system G x = g
GRAM_SOLVERS: dict[str, Callable[[torch.Tensor, torch.Tensor], torch.Tensor]] = {
"cholesky": _solve_cholesky,
"lu": _solve_lu,
"lstsq": _solve_lstsq,
}
[docs]
class EnergyNaturalGradientPreconditioner(MatrixPreconditionerPinn):
"""Energy-based natural gradient preconditioner.
Args:
space: The approximation space.
pde: The PDE to be solved, which can be an instance of EllipticPDE,
TemporalPDE, KineticPDE, or LinearOrder2PDE.
**kwargs: Additional keyword arguments.
Keyword Args:
matrix_regularization (float): Regularization parameter for the
preconditioning matrix (default: 1e-6).
`adaptive_matrix_regularization` (:code:`bool`, default=False):
If True, adaptively adjusts the regularization parameter during training
(à la Levenberg-Marquardt).
If False, uses the fixed regularization parameter
specified by `matrix_regularization`.
`adaptive_matrix_regularization_increase` (:code:`float`, default=10.0):
Factor by which to increase the regularization parameter if the line search
fails to find a suitable step size.
`adaptive_matrix_regularization_decrease` (:code:`float`, default=0.5):
Factor by which to decrease the regularization parameter if the line search
quickly succeeds in finding a suitable step size.
`adaptive_matrix_regularization_min` (:code:`float`, default=1e-12):
Minimum value for the regularization parameter when using adaptive adjustment.
gram_solver (str): dense solver for the natural-gradient system
``G @ d = grad``, one of ``GRAM_SOLVERS`` keys: ``"cholesky"``
(default; the regularized Gram is SPD, cheaper than a least-squares
solve), ``"lu"`` or ``"lstsq"`` (historical behavior).
use_lstsq (bool): Deprecated alias, kept for backward compatibility.
If given and `gram_solver` is not, ``True`` maps to
``gram_solver="lstsq"`` and ``False`` to ``gram_solver="lu"``
(default: unset, i.e. `gram_solver` applies).
gram_assembly (str): how to contract the per-point Jacobians ``Phi``
into the Gram matrix ``sum_i Phi_i @ Phi_i^T``. ``"gemm"`` (default):
single contraction over the batch and residual-component axes at
once (via :code:`torch.einsum`). ``"vmap"``: one matmul per residual
component (via :code:`torch.vmap`), then summed -- same result up to
round-off.
(historical) `gram_assembly`.
"""
def __init__(
self,
space: AbstractApproxSpace,
pde: ACCEPTED_PDE_TYPES,
**kwargs,
):
super().__init__(space, pde, **kwargs)
self.matrix_regularization = kwargs.get("matrix_regularization", 1e-6)
self.adaptive_matrix_regularization = kwargs.get(
"adaptive_matrix_regularization", False
)
self.adaptive_matrix_regularization_increase = kwargs.get(
"adaptive_matrix_regularization_increase", 10.0
)
self.adaptive_matrix_regularization_decrease = kwargs.get(
"adaptive_matrix_regularization_decrease", 0.5
)
self.adaptive_matrix_regularization_min = kwargs.get(
"adaptive_matrix_regularization_min", 1e-12
)
self.initial_matrix_regularization = self.matrix_regularization
gram_solver = kwargs.get("gram_solver")
if gram_solver is None:
# backward compatibility with the deprecated `use_lstsq` flag
use_lstsq = kwargs.get("use_lstsq")
if use_lstsq is None:
gram_solver = "cholesky"
else:
gram_solver = "lstsq" if use_lstsq else "lu"
self.gram_solver = gram_solver
assert self.gram_solver in GRAM_SOLVERS, (
f"unknown gram_solver {self.gram_solver!r}, "
f"expected one of {sorted(GRAM_SOLVERS)}"
)
self.gram_assembly = kwargs.get("gram_assembly", "gemm")
assert self.gram_assembly in ["vmap", "gemm"]
def _gram_contraction(self, Phi: torch.Tensor) -> torch.Tensor: # noqa: N803
"""Contract per-point Jacobians into a Gram matrix.
Computes ``sum_{i,s} Phi[i, :, s] outer Phi[i, :, s]``, i.e. the sum over
collocation points (axis 0) and residual components (axis 2) of the outer
product of the ``ndof``-sized Jacobian slices (axis 1).
Args:
Phi: Per-point Jacobians, of shape ``(N, ndof, size)``.
Returns:
The ``(ndof, ndof)`` Gram matrix contribution.
"""
if self.gram_assembly == "vmap":
vmap = torch.vmap(lambda mat: mat.T @ mat, in_dims=2, out_dims=2)
per_component = vmap(Phi)
return per_component.sum(dim=-1)
return torch.einsum("ijk,ilk->jl", Phi, Phi)
[docs]
def compute_preconditioning_matrix(
self, labels: torch.Tensor, *args: torch.Tensor, **kwargs
) -> torch.Tensor:
"""Compute the preconditioning matrix using the main operator.
Args:
labels: The labels tensor.
*args: Additional arguments.
**kwargs: Additional keyword arguments.
Returns:
The preconditioning matrix.
"""
with torch.no_grad():
N = args[0].shape[0]
theta = self.get_formatted_current_theta()
Phi = self.operator.apply_dict_of_vmap_to_label_tensors(
self.vectorized_Phi, theta, labels, *args
)
if len(self.in_weights) == 1: # apply the same weights to all labels
for key in self.in_weights: # dummy loop
Phi[:, :, :] *= self.in_weights[key]
else: # apply weights for each labels
for key in self.in_weights:
Phi[labels == key, :, :] *= self.in_weights[key]
M = self._gram_contraction(Phi) / N
M.diagonal().add_(self.matrix_regularization)
return self.in_weight * M
[docs]
def compute_preconditioning_matrix_bc(
self, labels: torch.Tensor, *args: torch.Tensor, **kwargs
) -> torch.Tensor:
"""Compute the boundary condition preconditioning matrix.
Args:
labels: The labels tensor.
*args: Additional arguments.
**kwargs: Additional keyword arguments.
Returns:
The boundary condition preconditioning matrix.
"""
with torch.no_grad():
N = args[0].shape[0]
theta = self.get_formatted_current_theta()
self.operator_bc = cast(FunctionalOperator, self.operator_bc)
self.vectorized_Phi_bc = cast(TYPE_DICT_OF_VMAPS, self.vectorized_Phi_bc)
Phi = self.operator_bc.apply_dict_of_vmap_to_label_tensors(
self.vectorized_Phi_bc, theta, labels, *args
)
if len(self.bc_weights) == 1: # apply the same weights to all labels
for key in self.bc_weights: # dummy loop
Phi[:, :, :] *= self.bc_weights[key]
else: # apply weights for each labels
for key in self.bc_weights:
Phi[labels == key, :, :] *= self.bc_weights[key]
M = self._gram_contraction(Phi) / N
M.diagonal().add_(self.matrix_regularization)
return self.bc_weight * M
[docs]
def compute_preconditioning_matrix_ic(
self, labels: torch.Tensor, *args: torch.Tensor, **kwargs
) -> torch.Tensor:
"""Compute the initial condition preconditioning matrix.
Args:
labels: The labels tensor.
*args: Additional arguments.
**kwargs: Additional keyword arguments.
Returns:
The initial condition preconditioning matrix.
"""
with torch.no_grad():
N = args[0].shape[0]
theta = self.get_formatted_current_theta()
self.operator_ic = cast(FunctionalOperator, self.operator_ic)
self.vectorized_Phi_ic = cast(TYPE_DICT_OF_VMAPS, self.vectorized_Phi_ic)
Phi = self.operator_ic.apply_dict_of_vmap_to_label_tensors(
self.vectorized_Phi_ic, theta, labels, *args
)
if len(self.ic_weights) == 1: # apply the same weights to all labels
for key in self.ic_weights: # dummy loop
Phi[:, :, :] *= self.ic_weights[key]
else: # apply weights for each labels
for key in self.ic_weights:
Phi[labels == key, :, :] *= self.ic_weights[key]
M = self._gram_contraction(Phi) / N
M.diagonal().add_(self.matrix_regularization)
return self.ic_weight * M
[docs]
def compute_preconditioning_matrix_dl(
self, *args: torch.Tensor, **kwargs
) -> torch.Tensor:
"""Computes the Gram matrix of the network for the given input tensors.
Args:
*args: Input tensors for computing the Gram matrix.
**kwargs: Additional keyword arguments.
Returns:
The computed Gram matrix.
"""
with torch.no_grad():
N = args[0].shape[0]
jacobian = self.space.jacobian(*args)
M = self._gram_contraction(jacobian) / N
M.diagonal().add_(self.matrix_regularization)
return M
[docs]
def compute_full_preconditioning_matrix(
self, data: tuple | dict, **kwargs
) -> torch.Tensor:
"""Compute the full preconditioning matrix.
Include contributions from the main operator, boundary conditions, and initial
conditions.
Args:
data: Input data, either as a tuple or a dictionary.
**kwargs: Additional keyword arguments.
Returns:
The full preconditioning matrix.
"""
M = self.get_preconditioning_matrix(data, **kwargs)
if self.has_bc:
M += self.get_preconditioning_matrix_bc(data, **kwargs)
if self.has_ic:
M += self.get_preconditioning_matrix_ic(data, **kwargs)
for index, coeff in enumerate(self.dl_weights):
M += coeff * self.get_preconditioning_matrix_dl(
self.args_for_dl[index], **kwargs
)
return M
def __call__(
self,
epoch: int,
data: tuple | dict,
grads: torch.Tensor,
res_l: tuple,
res_r: tuple,
**kwargs,
) -> torch.Tensor:
"""Apply the preconditioner to the input gradients.
Args:
epoch: Current training epoch.
data: Input data, either as a tuple or a dictionary.
grads: Gradient tensor to be preconditioned.
res_l: Left residuals.
res_r: Right residuals.
**kwargs: Additional keyword arguments.
Returns:
The preconditioned gradient tensor.
"""
with torch.no_grad():
M = self.compute_full_preconditioning_matrix(data, **kwargs)
return GRAM_SOLVERS[self.gram_solver](M, grads)
[docs]
def update_matrix_regularization(self, n_steps: int, loss_has_decreased: bool):
"""Updates the regularization parameter for the preconditioning matrix.
This method adaptively adjusts the regularization parameter during training
based on the number of line search steps taken
and whether the loss has decreased.
Args:
n_steps: The number of line search steps taken.
loss_has_decreased: Whether the loss has decreased after the line search.
"""
if self.adaptive_matrix_regularization:
if loss_has_decreased:
if n_steps <= 1:
self.matrix_regularization *= (
self.adaptive_matrix_regularization_decrease
)
else:
self.matrix_regularization *= (
self.adaptive_matrix_regularization_increase
)
self.matrix_regularization = max(
self.matrix_regularization, self.adaptive_matrix_regularization_min
)
[docs]
def reset_matrix_regularization(self):
"""Resets the regularization parameter to its initial value."""
self.matrix_regularization = self.initial_matrix_regularization