Accelerating Natural Gradient Descent for PINNs with Randomized Numerical Linear Algebra
Authors: Ivan Bioli, Carlo Marcati, Giancarlo Sangalli
Organizations: Department of Mathematics, University of Pavia, Via A. Ferrata 5, 27100 Pavia, Italy · Department of Civil Engineering and Architecture, University of Pavia, Via A. Ferrata 3, 27100 Pavia, Italy · Istituto di Matematica Applicata e Tecnologie Informatiche “E. Magenes”, CNR, Via A. Ferrata 1, 27100 Pavia, Italy
Natural Gradient Descent (NGD) has emerged as a promising optimization algorithm for training neural network-based solvers for partial differential equations (PDEs), such as Physics-Informed Neural Networks (PINNs). However, its practical use is often limited by the high computational cost of solving linear systems involving the Gramian matrix. While matrix-free NGD methods based on the conjugate gradient (CG) method avoid explicit matrix inversion, the ill-conditioning of the Gramian significantly slows the convergence of CG. In this work, we extend matrix-free NGD to broader classes of problems than previously considered and propose the use of Randomized Numerical Linear Algebra (RandNLA) techniques for efficient preconditioning of the inner CG solver. The resulting algorithms demonstrate substantial performance improvements over existing NGD-based methods on a range of PDE problems discretized using neural networks, and offer competitive results compared to other state-of-the-art optimizers.
Figures & tables
Strategy
Setup Cost
MMP Cost
Memory Footprint
Precompute Gramian
O(q[FP]+p2q)
O(p2ℓ)
O(p2+pB)
Precompute Jacobian only
O(q[FP])
O(pqℓ)
O(pq)
Recompute batch Jacobian
on the fly
–
O(q[FP]+pqℓ)
O(pB)
Table 1 . Summary of computational and memory costs for the three matrix-vector product strategies. Here p is the number of parameters, [FP] is the cost of a forward pass, ℓ is the number of simultaneous matvecs, q is the number of quadrature points, and B is the batch size for processing quadrature points. We assume that p≤q , i.e., we focus on the regime where a highly accurate quadrature rule is employed and the number of quadrature points q exceeds the number of neural network parameters p .
Optimizer
Poisson 3D
Heat 3+1D
Kovasznay
Deep-Ritz Poisson
FEINNs 2D
NGD full
1.56×10−6
(1.37−1.77×10−6)
317.7 s
1.75×10−6
(1.49−1.86×10−6)
279.2 s
1.08×10−6
(0.965−1.25×10−6)
Table 2 . Relative H1(Ω) error at the plateau iterate for each optimizer and test problem. The first line reports the median error, the second line reports in smaller font and in parentheses the 25 th and 75 th percentiles in compact form as (a−b×10k) , and the third line reports the corresponding cumulative wall-clock time in seconds. The plateau criterion and the dot ( ∙ ) marking the corresponding iterate in the convergence curves are described in Section 5.2 ; see Figure 10 .
Optimizer
3×64×64×64×1
q=104
3×64×64×64×1
q=105
3×128×128×128×1
q=105
NGD full
1.56×10−6
(1.37−1.77×10−6)
317.7 s
Table 3 . Scalability experiments on the 3D Poisson problem discretized via PINNs; see Section 5.3 . Relative H1(Ω) error at the plateau iterate for each optimizer and architecture. The first line reports the median error, the second line reports in smaller font and in parentheses the 25 th and 75 th percentiles in compact form as (a−b×10k) , and the third line reports the corresponding cumulative wall-clock time in seconds, while the fourth line reports the per-iteration cost in seconds. Each column corresponds to a different network architecture and number of quadrature points q , as indicated in the header.
Appendix figures & tables6 assets
Supplementary material from the paper’s appendix.
Appendix
Experiment
NGD (all)
(L-)BFGS
SSBroyden
Adam
Network Architecture
Poisson 3D
1 000
50 000
50 000
150 000
[3, 64, 64, 64, 1]
Heat 3+1D
1 000
50 000
50 000
150 000
[4, 64, 64, 64, 1]
Kovasznay
2 000
25 000
25 000
75 000
u : [2, 45, 45, 45, 2]
p : [2, 45, 45, 45, 1]
Deep-Ritz Poisson
1 000
25 000
25 000
75 000
[2, 64, 64, 64, 1]
FEINNs 2D
1 000
25 000
25 000
75 000
[2, 64, 64, 64, 1]
Appendix
Table 4 . Iteration budgets for the evaluated optimizers and neural network architectures for each PDE experiment. Network shapes are denoted as [input_dim, hidden_1, ..., hidden_N, output_dim] .
Experiment
Interior / Mesh configuration
Boundary Points
Evaluation set
Poisson 3D
10 000 random
1 000 random
100 000 random
Heat 3+1D
10 000 random
1 000 random
100 000 random
Kovasznay
10 000 random
1 000 random
100 000 random
Deep-Ritz Poisson
Quad mesh, h=1/30 , d=4 (14.5k DoF)
-
Gauss–Jacobi, order 16
FEINNs 2D
Quad mesh, h=1/30 , d=4 (14.5k DoF)
-
Gauss–Jacobi, order 16
Appendix
Table 5 . Domain discretization, boundary conditions, and evaluation setups for each numerical experiment.
γ
Poisson 3D
Heat 3+1D
Deep-Ritz Poisson
FEINNs 2D
0.1
1.87×10−6
(0.886−3.27×10−6)
1.004×
7.42×10−6
(6.05−8.60×10−6)
1.023×
7.49×10−1
(3.99−10.4×10−1)
Appendix
Table 6 . Sensitivity analysis of γ for NyströmNGD with Gaussian test matrices. The first line reports the median relative H1(Ω) error, the second line reports in smaller font and in parentheses the 25 th and 75 th percentiles in compact form as (a−b×10k) , and the third line reports the wall-clock time relative to the default γ=10 .
γ
Poisson 3D
Heat 3+1D
Deep-Ritz Poisson
FEINNs 2D
0.1
8.42×10−3
(7.79−8.49×10−3)
1.062×
7.48×10−3
(5.24−9.63×10−3)
1.055×
2.12
(1.08−2.53)
Appendix
Table 7 . Sensitivity analysis of γ for RPCholNGD. The first line reports the median relative H1(Ω) error, the second line reports in smaller font and in parentheses the 25 th and 75 th percentiles in compact form as (a−b×10k) , and the third line reports the wall-clock time relative to the default γ=10 .
r
Poisson 3D
Heat 3+1D
Kovasznay
Deep-Ritz Poisson
FEINNs 2D
0.1
1.68×10−6
(1.55−1.75×10−6)
1.00×
2.03×10−6
(2.03−2.12×10−6)
1.00×
6.95×10−7
(6.80−7.48×10−7)
Appendix
Table 8 . Sensitivity analysis of the parameter r for NyströmNGD with Gaussian test matrices. The first line reports the median relative H1(Ω) error, the second line reports in smaller font and in parentheses the 25 th and 75 th percentiles in compact form as (a−b×10k) , and the third line reports the wall-clock time relative to the default r=10 .
r
Poisson 3D
Heat 3+1D
Kovasznay
Deep-Ritz Poisson
FEINNs 2D
0.01
1.83×10−6
(1.78−1.87×10−6)
1.00×
2.11×10−6
(1.81−2.53×10−6)
1.00×
9.79×10−7
(9.14−10.3×10−7)
Appendix
Table 9 . Sensitivity analysis of the RPCholesky stopping tolerance scaling factor r ( ε=r⋅p⋅ϵmach ) for RPCholNGD. The first line reports the median relative H1(Ω) error, the second line reports in smaller font and in parentheses the 25 th and 75 th percentiles in compact form as (a−b×10k) , and the third line reports the wall-clock time relative to the default r=1 .
National Center for Applied Mathematics Tianjin University Tianjin, 300072, China · School of Mechanical and Aerospace Engineering Jilin University Changchun, 130025, China
Department of Information Science and Engineering, KTH Royal Institute of Technology, Stockholm, Sweden · School of Advanced Manufacturing and Robotics, Peking University, Beijing, China · School of Advanced Technology, Xi’an Jiaotong-Liverpool University, Suzhou, China +2