Source code for scimba_torch.numerical_solvers.pinn_preconditioners.energy_ng

"""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