From 7cbc6db9fd42bbf125ed16ce285bf87345e6bd65 Mon Sep 17 00:00:00 2001 From: Laurent Hartwich Date: Tue, 14 Jul 2026 11:50:48 +0200 Subject: [PATCH 1/4] timeit wrapper: ensure kwargs has key :timer before checking its value --- src/NaturalGradient/NaturalGradient.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/NaturalGradient/NaturalGradient.jl b/src/NaturalGradient/NaturalGradient.jl index 4e118d1..9494812 100644 --- a/src/NaturalGradient/NaturalGradient.jl +++ b/src/NaturalGradient/NaturalGradient.jl @@ -173,7 +173,7 @@ function tdvp_relative_error(J::Jacobian, Es::EnergySummary, θdot::Vector) end function NaturalGradient_timeit_wrapper(θ, Oks_and_Eks_; kwargs...) - if kwargs[:timer] !== nothing + if haskey(kwargs, :timer) && kwargs[:timer] !== nothing ng = @timeit kwargs[:timer] "NaturalGradient" NaturalGradient(θ, Oks_and_Eks_; kwargs...) else ng = NaturalGradient(θ, Oks_and_Eks_; kwargs...) From 785de363371a4bfce7d1025a76264739b3ee0d58 Mon Sep 17 00:00:00 2001 From: Laurent Hartwich Date: Tue, 14 Jul 2026 14:20:59 +0200 Subject: [PATCH 2/4] add docstring --- src/solver/eigen_solver.jl | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/src/solver/eigen_solver.jl b/src/solver/eigen_solver.jl index cd189e4..97801ea 100644 --- a/src/solver/eigen_solver.jl +++ b/src/solver/eigen_solver.jl @@ -8,6 +8,11 @@ mutable struct EigenSolver <: AbstractSolver end +""" +(solver::AbstractSolver)(M::AbstractMatrix, v::AbstractArray, double::Bool; method=:auto, kwargs...) + +returns o = M^-1 v (computed by eigendecomposition) +""" function (solver::EigenSolver)(M::AbstractMatrix, v::AbstractArray) #@assert ishermitian(M) "EigenSolver: M is not Hermitian" eig = eigen(Hermitian(M)) From 349923be31ff381f045e432794237924c53b51e2 Mon Sep 17 00:00:00 2001 From: Laurent Hartwich Date: Tue, 14 Jul 2026 14:21:54 +0200 Subject: [PATCH 3/4] ensure LinearSolve wrapper works when (only) one of the arguments is complex --- src/solver/LinearSolveWrapper.jl | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/src/solver/LinearSolveWrapper.jl b/src/solver/LinearSolveWrapper.jl index 1662f47..5050e20 100644 --- a/src/solver/LinearSolveWrapper.jl +++ b/src/solver/LinearSolveWrapper.jl @@ -17,6 +17,12 @@ function (solver::LinearSolveWrapper)(M::AbstractMatrix, v::AbstractArray) else M = Hermitian(M) end + + if eltype(M) != eltype(v) + if eltype(M) === ComplexF64 + v = ComplexF64.(v) + end + end prob = LinearProblem(M, v) sol = solve(prob, solver.alg) From 1fa76ea9f41547cbe1cd9eef46cf472431eb4332 Mon Sep 17 00:00:00 2001 From: Laurent Hartwich Date: Tue, 14 Jul 2026 14:22:28 +0200 Subject: [PATCH 4/4] add docstrings and timers --- src/solver/solver.jl | 28 +++++++++++++++++++++------- 1 file changed, 21 insertions(+), 7 deletions(-) diff --git a/src/solver/solver.jl b/src/solver/solver.jl index 3997913..5f95dd1 100644 --- a/src/solver/solver.jl +++ b/src/solver/solver.jl @@ -1,27 +1,41 @@ abstract type AbstractSolver end +""" +solve_S(solver::AbstractSolver, J::Jacobian, grad_half::Vector; timer=TimerOutput(), kwargs...) + +In the terms of https://arxiv.org/pdf/2503.12557, this returns θdot = (O' * O)^-1 * O' * E_loc (Eq. 11) +""" function solve_S(solver::AbstractSolver, J::Jacobian, grad_half::Vector; timer=TimerOutput(), kwargs...) - @timeit "dense_S" Jd = dense_S(J) - @timeit "solve" θdot = -solver(Jd, grad_half; kwargs...) + @timeit timer "dense_S" Jd = dense_S(J) + @timeit timer "solve" θdot = -solver(Jd, grad_half; kwargs...) return θdot end +""" +solve_T(solver::AbstractSolver, J::Jacobian, Es::EnergySummary; timer=TimerOutput(), kwargs...) + +In the terms of https://arxiv.org/pdf/2503.12557, this returns θdot = O' * (O * O')^-1 * E_loc (Eq. 13) +""" function solve_T(solver::AbstractSolver, J::Jacobian, Es::EnergySummary; timer=TimerOutput(), kwargs...) - @timeit "dense_T" Jd = dense_T(J) + @timeit timer "dense_T" Jd = dense_T(J) Ekms = centered(Es) - - @timeit "solve" θdot_raw = -solver(Jd, Ekms; kwargs...) - θdot = centered(J)' * θdot_raw + @timeit timer "solve" θdot_raw = -solver(Jd, Ekms; kwargs...) + @timeit timer "mult." θdot = centered(J)' * θdot_raw return θdot end +""" +(solver::AbstractSolver)(ng::NaturalGradient; method=:auto, compute_error=true, kwargs...) + +Computes θdot corresponding to equations (11) or (13) of https://arxiv.org/pdf/2503.12557, depending on the method (solve_T corresponds to eq. (13), while solve_S corresponds to eq. (11)). +""" function (solver::AbstractSolver)(ng::NaturalGradient; method=:auto, compute_error=true, kwargs...) if method === :T || (method === :auto && nr_samples(ng.J) < nr_parameters(ng.J)) ng.θdot = solve_T(solver, ng.J, ng.Es; kwargs...) else - ng.θdot = solve_S(solver, ng.J, get_gradient(ng) ./ 2; kwargs...) + ng.θdot = solve_S(solver, ng.J, get_gradient_timeit_wrapper(ng; kwargs...) ./ 2; kwargs...) end if compute_error tdvp_error!(ng)