diff --git a/Project.toml b/Project.toml index d55932ba7..3990f7c30 100644 --- a/Project.toml +++ b/Project.toml @@ -213,6 +213,7 @@ InteractiveUtils = "b77e0a4c-d291-57a0-90e8-8db25a27a240" IterativeSolvers = "42fd0dbc-a981-5370-80f2-aaf504508153" JacobiDavidson = "11c68b98-9c9b-11e8-267b-bbb95576cead" KrylovKit = "0b1a1467-8014-51b9-945f-bf0ae24f4b77" +LAPACK_jll = "51474c39-65e3-53ba-86ba-03b1b862ec14" MultiFloats = "bdf0d083-296b-4888-a5b6-7498122e68a5" Pkg = "44cfe95a-1eb2-52ea-b672-e2afdf69b78f" PureUMFPACK = "b7e1f0a2-3c4d-4e5f-9a0b-1c2d3e4f5a6b" @@ -228,6 +229,7 @@ StableRNGs = "860ef19b-820b-49d6-a774-d7a799459cd3" StaticArrays = "90137ffa-7385-5640-81b9-e52037218182" Test = "8dfed614-e22c-5e08-85e1-65c5234f0b40" Zygote = "e88e6eb3-aa80-5325-afca-941959d7151f" +blis_jll = "6136c539-28a5-5bf0-87cc-b183200dce32" [targets] -test = ["AlgebraicMultigrid", "Arpack", "ArnoldiMethod", "BandedMatrices", "BlockDiagonals", "CliqueTrees", "ComponentArrays", "FastAlmostBandedMatrices", "FastLapackInterface", "FiniteDiff", "FixedSizeArrays", "ForwardDiff", "InteractiveUtils", "IterativeSolvers", "JacobiDavidson", "KrylovKit", "MultiFloats", "Pkg", "PureUMFPACK", "Random", "RecursiveFactorization", "STRUMPACK_jll", "SafeTestsets", "SciMLTesting", "SparseArrays", "Sparspak", "SpecializingFactorizations", "StableRNGs", "StaticArrays", "Test", "Zygote"] +test = ["AlgebraicMultigrid", "Arpack", "ArnoldiMethod", "BandedMatrices", "BlockDiagonals", "CliqueTrees", "ComponentArrays", "FastAlmostBandedMatrices", "FastLapackInterface", "FiniteDiff", "FixedSizeArrays", "ForwardDiff", "InteractiveUtils", "IterativeSolvers", "JacobiDavidson", "KrylovKit", "LAPACK_jll", "MultiFloats", "Pkg", "PureUMFPACK", "Random", "RecursiveFactorization", "STRUMPACK_jll", "SafeTestsets", "SciMLTesting", "SparseArrays", "Sparspak", "SpecializingFactorizations", "StableRNGs", "StaticArrays", "Test", "Zygote", "blis_jll"] diff --git a/ext/LinearSolveBLISExt.jl b/ext/LinearSolveBLISExt.jl index a290cbad0..0bf84087c 100644 --- a/ext/LinearSolveBLISExt.jl +++ b/ext/LinearSolveBLISExt.jl @@ -6,34 +6,65 @@ using LAPACK_jll using LinearAlgebra using LinearSolve -using LinearAlgebra: BlasInt, LU +using LinearAlgebra: BlasInt using LinearAlgebra.LAPACK: require_one_based_indexing, chkfinite, chkstride1, @blasfunc, chkargsok -using LinearSolve: ArrayInterface, BLISLUFactorization, @get_cacheval, LinearCache, SciMLBase, LinearVerbosity, get_blas_operation_info, blas_info_msg +using LinearSolve: BLISLUFactorization, @get_cacheval, LinearCache, SciMLBase, LinearVerbosity, get_blas_operation_info, blas_info_msg using SciMLLogging: SciMLLogging, @SciMLMessage using SciMLBase: ReturnCode const global libblis = blis_jll.blis const global liblapack = LAPACK_jll.liblapack +# Resolve Julia 1.13 lazy JLL products once so solves call fixed function pointers. +const _lapack_handle = Ref{Ptr{Cvoid}}(C_NULL) +const _lapack_zgetrf = Ref{Ptr{Cvoid}}(C_NULL) +const _lapack_cgetrf = Ref{Ptr{Cvoid}}(C_NULL) +const _lapack_dgetrf = Ref{Ptr{Cvoid}}(C_NULL) +const _lapack_sgetrf = Ref{Ptr{Cvoid}}(C_NULL) +const _lapack_zgetrs = Ref{Ptr{Cvoid}}(C_NULL) +const _lapack_cgetrs = Ref{Ptr{Cvoid}}(C_NULL) +const _lapack_dgetrs = Ref{Ptr{Cvoid}}(C_NULL) +const _lapack_sgetrs = Ref{Ptr{Cvoid}}(C_NULL) + +function __init__() + @static if VERSION >= v"1.13.0-DEV.0" + handle = Libdl.dlopen(liblapack) + _lapack_handle[] = handle + _lapack_zgetrf[] = Libdl.dlsym(handle, @blasfunc(zgetrf_)) + _lapack_cgetrf[] = Libdl.dlsym(handle, @blasfunc(cgetrf_)) + _lapack_dgetrf[] = Libdl.dlsym(handle, @blasfunc(dgetrf_)) + _lapack_sgetrf[] = Libdl.dlsym(handle, @blasfunc(sgetrf_)) + _lapack_zgetrs[] = Libdl.dlsym(handle, @blasfunc(zgetrs_)) + _lapack_cgetrs[] = Libdl.dlsym(handle, @blasfunc(cgetrs_)) + _lapack_dgetrs[] = Libdl.dlsym(handle, @blasfunc(dgetrs_)) + _lapack_sgetrs[] = Libdl.dlsym(handle, @blasfunc(sgetrs_)) + end + return nothing +end + +macro _lapack_function(symbol, pointer) + if VERSION >= v"1.13.0-DEV.0" + return :($(esc(pointer))[]) + end + return :(($(esc(symbol)), liblapack)) +end + LinearSolve.useblis(x::Nothing) = true -function getrf!( - A::AbstractMatrix{<:ComplexF64}; - ipiv = similar(A, BlasInt, min(size(A, 1), size(A, 2))), - info = Ref{BlasInt}(), - check = false +@inline function getrf!( + A::AbstractMatrix{<:ComplexF64}, ipiv::AbstractVector{BlasInt}, + info::Ref{BlasInt}, check::Bool ) require_one_based_indexing(A) check && chkfinite(A) chkstride1(A) m, n = size(A) lda = max(1, stride(A, 2)) - if isempty(ipiv) - ipiv = similar(A, BlasInt, min(size(A, 1), size(A, 2))) - end + length(ipiv) == min(m, n) || + throw(DimensionMismatch("ipiv has length $(length(ipiv)), but needs $(min(m, n))")) ccall( - (@blasfunc(zgetrf_), liblapack), Cvoid, + @_lapack_function(@blasfunc(zgetrf_), _lapack_zgetrf), Cvoid, ( Ref{BlasInt}, Ref{BlasInt}, Ptr{ComplexF64}, Ref{BlasInt}, Ptr{BlasInt}, Ptr{BlasInt}, @@ -41,25 +72,22 @@ function getrf!( m, n, A, lda, ipiv, info ) chkargsok(info[]) - return A, ipiv, info[], info #Error code is stored in LU factorization type + return info[] end -function getrf!( - A::AbstractMatrix{<:ComplexF32}; - ipiv = similar(A, BlasInt, min(size(A, 1), size(A, 2))), - info = Ref{BlasInt}(), - check = false +@inline function getrf!( + A::AbstractMatrix{<:ComplexF32}, ipiv::AbstractVector{BlasInt}, + info::Ref{BlasInt}, check::Bool ) require_one_based_indexing(A) check && chkfinite(A) chkstride1(A) m, n = size(A) lda = max(1, stride(A, 2)) - if isempty(ipiv) - ipiv = similar(A, BlasInt, min(size(A, 1), size(A, 2))) - end + length(ipiv) == min(m, n) || + throw(DimensionMismatch("ipiv has length $(length(ipiv)), but needs $(min(m, n))")) ccall( - (@blasfunc(cgetrf_), liblapack), Cvoid, + @_lapack_function(@blasfunc(cgetrf_), _lapack_cgetrf), Cvoid, ( Ref{BlasInt}, Ref{BlasInt}, Ptr{ComplexF32}, Ref{BlasInt}, Ptr{BlasInt}, Ptr{BlasInt}, @@ -67,25 +95,22 @@ function getrf!( m, n, A, lda, ipiv, info ) chkargsok(info[]) - return A, ipiv, info[], info #Error code is stored in LU factorization type + return info[] end -function getrf!( - A::AbstractMatrix{<:Float64}; - ipiv = similar(A, BlasInt, min(size(A, 1), size(A, 2))), - info = Ref{BlasInt}(), - check = false +@inline function getrf!( + A::AbstractMatrix{<:Float64}, ipiv::AbstractVector{BlasInt}, + info::Ref{BlasInt}, check::Bool ) require_one_based_indexing(A) check && chkfinite(A) chkstride1(A) m, n = size(A) lda = max(1, stride(A, 2)) - if isempty(ipiv) - ipiv = similar(A, BlasInt, min(size(A, 1), size(A, 2))) - end + length(ipiv) == min(m, n) || + throw(DimensionMismatch("ipiv has length $(length(ipiv)), but needs $(min(m, n))")) ccall( - (@blasfunc(dgetrf_), liblapack), Cvoid, + @_lapack_function(@blasfunc(dgetrf_), _lapack_dgetrf), Cvoid, ( Ref{BlasInt}, Ref{BlasInt}, Ptr{Float64}, Ref{BlasInt}, Ptr{BlasInt}, Ptr{BlasInt}, @@ -93,25 +118,22 @@ function getrf!( m, n, A, lda, ipiv, info ) chkargsok(info[]) - return A, ipiv, info[], info #Error code is stored in LU factorization type + return info[] end -function getrf!( - A::AbstractMatrix{<:Float32}; - ipiv = similar(A, BlasInt, min(size(A, 1), size(A, 2))), - info = Ref{BlasInt}(), - check = false +@inline function getrf!( + A::AbstractMatrix{<:Float32}, ipiv::AbstractVector{BlasInt}, + info::Ref{BlasInt}, check::Bool ) require_one_based_indexing(A) check && chkfinite(A) chkstride1(A) m, n = size(A) lda = max(1, stride(A, 2)) - if isempty(ipiv) - ipiv = similar(A, BlasInt, min(size(A, 1), size(A, 2))) - end + length(ipiv) == min(m, n) || + throw(DimensionMismatch("ipiv has length $(length(ipiv)), but needs $(min(m, n))")) ccall( - (@blasfunc(sgetrf_), liblapack), Cvoid, + @_lapack_function(@blasfunc(sgetrf_), _lapack_sgetrf), Cvoid, ( Ref{BlasInt}, Ref{BlasInt}, Ptr{Float32}, Ref{BlasInt}, Ptr{BlasInt}, Ptr{BlasInt}, @@ -119,15 +141,15 @@ function getrf!( m, n, A, lda, ipiv, info ) chkargsok(info[]) - return A, ipiv, info[], info #Error code is stored in LU factorization type + return info[] end function getrs!( trans::AbstractChar, A::AbstractMatrix{<:ComplexF64}, ipiv::AbstractVector{BlasInt}, - B::AbstractVecOrMat{<:ComplexF64}; - info = Ref{BlasInt}() + B::AbstractVecOrMat{<:ComplexF64}, + info::Ref{BlasInt} ) require_one_based_indexing(A, ipiv, B) LinearAlgebra.LAPACK.chktrans(trans) @@ -141,7 +163,7 @@ function getrs!( end nrhs = size(B, 2) ccall( - (@blasfunc(zgetrs_), liblapack), Cvoid, + @_lapack_function(@blasfunc(zgetrs_), _lapack_zgetrs), Cvoid, ( Ref{UInt8}, Ref{BlasInt}, Ref{BlasInt}, Ptr{ComplexF64}, Ref{BlasInt}, Ptr{BlasInt}, Ptr{ComplexF64}, Ref{BlasInt}, Ptr{BlasInt}, Clong, @@ -157,8 +179,8 @@ function getrs!( trans::AbstractChar, A::AbstractMatrix{<:ComplexF32}, ipiv::AbstractVector{BlasInt}, - B::AbstractVecOrMat{<:ComplexF32}; - info = Ref{BlasInt}() + B::AbstractVecOrMat{<:ComplexF32}, + info::Ref{BlasInt} ) require_one_based_indexing(A, ipiv, B) LinearAlgebra.LAPACK.chktrans(trans) @@ -172,7 +194,7 @@ function getrs!( end nrhs = size(B, 2) ccall( - (@blasfunc(cgetrs_), liblapack), Cvoid, + @_lapack_function(@blasfunc(cgetrs_), _lapack_cgetrs), Cvoid, ( Ref{UInt8}, Ref{BlasInt}, Ref{BlasInt}, Ptr{ComplexF32}, Ref{BlasInt}, Ptr{BlasInt}, Ptr{ComplexF32}, Ref{BlasInt}, Ptr{BlasInt}, Clong, @@ -188,8 +210,8 @@ function getrs!( trans::AbstractChar, A::AbstractMatrix{<:Float64}, ipiv::AbstractVector{BlasInt}, - B::AbstractVecOrMat{<:Float64}; - info = Ref{BlasInt}() + B::AbstractVecOrMat{<:Float64}, + info::Ref{BlasInt} ) require_one_based_indexing(A, ipiv, B) LinearAlgebra.LAPACK.chktrans(trans) @@ -203,7 +225,7 @@ function getrs!( end nrhs = size(B, 2) ccall( - (@blasfunc(dgetrs_), liblapack), Cvoid, + @_lapack_function(@blasfunc(dgetrs_), _lapack_dgetrs), Cvoid, ( Ref{UInt8}, Ref{BlasInt}, Ref{BlasInt}, Ptr{Float64}, Ref{BlasInt}, Ptr{BlasInt}, Ptr{Float64}, Ref{BlasInt}, Ptr{BlasInt}, Clong, @@ -219,8 +241,8 @@ function getrs!( trans::AbstractChar, A::AbstractMatrix{<:Float32}, ipiv::AbstractVector{BlasInt}, - B::AbstractVecOrMat{<:Float32}; - info = Ref{BlasInt}() + B::AbstractVecOrMat{<:Float32}, + info::Ref{BlasInt} ) require_one_based_indexing(A, ipiv, B) LinearAlgebra.LAPACK.chktrans(trans) @@ -234,7 +256,7 @@ function getrs!( end nrhs = size(B, 2) ccall( - (@blasfunc(sgetrs_), liblapack), Cvoid, + @_lapack_function(@blasfunc(sgetrs_), _lapack_sgetrs), Cvoid, ( Ref{UInt8}, Ref{BlasInt}, Ref{BlasInt}, Ptr{Float32}, Ref{BlasInt}, Ptr{BlasInt}, Ptr{Float32}, Ref{BlasInt}, Ptr{BlasInt}, Clong, @@ -249,17 +271,28 @@ end default_alias_A(::BLISLUFactorization, ::Any, ::Any) = false default_alias_b(::BLISLUFactorization, ::Any, ::Any) = false -const PREALLOCATED_BLIS_LU = begin - A = rand(0, 0) - luinst = ArrayInterface.lu_instance(A), Ref{BlasInt}() +mutable struct BLISLUCache{F, P, I} + factors::F + ipiv::P + info::I end -function LinearSolve.init_cacheval( - alg::BLISLUFactorization, A::Matrix{Float64}, b, u, Pl, Pr, - maxiters::Int, abstol, reltol, verbose::Union{LinearVerbosity, Bool}, - assumptions::OperatorAssumptions +LinearSolve._custom_cache_factorization(::BLISLUFactorization, cacheval::BLISLUCache) = + LinearAlgebra.LU(cacheval.factors, cacheval.ipiv, Int(cacheval.info[])) + +@inline function LinearSolve._direct_lu_factorize!( + cacheval::BLISLUCache, A_work, ::BLISLUFactorization + ) + cacheval.factors = A_work + return getrf!(A_work, cacheval.ipiv, cacheval.info, false) +end + +@inline function LinearSolve._direct_lu_solve!( + cacheval::BLISLUCache, u, b, ::BLISLUFactorization ) - return PREALLOCATED_BLIS_LU + copyto!(u, b) + getrs!('N', cacheval.factors, cacheval.ipiv, u, cacheval.info) + return u end function LinearSolve.init_cacheval( @@ -268,60 +301,67 @@ function LinearSolve.init_cacheval( maxiters::Int, abstol, reltol, verbose::Union{LinearVerbosity, Bool}, assumptions::OperatorAssumptions ) - # Ask `lu_instance` about `A` itself so wrapper-specific dispatch can choose - # the factorization container that `solve!` will store. - return ArrayInterface.lu_instance(A), Ref{BlasInt}() + return BLISLUCache( + A, similar(A, BlasInt, min(size(A, 1), size(A, 2))), Ref{BlasInt}() + ) end function SciMLBase.solve!( cache::LinearCache, alg::BLISLUFactorization; kwargs... ) - A = cache.A - A = convert(AbstractMatrix, A) + A_work = convert(AbstractMatrix, cache.A) verbose = cache.verbose if cache.isfresh cacheval = @get_cacheval(cache, :BLISLUFactorization) - res = getrf!(A; ipiv = cacheval[1].ipiv, info = cacheval[2]) - fact = LU(res[1:3]...), res[4] - cache.cacheval = fact - - info_value = res[3] + if length(cacheval.ipiv) != min(size(A_work, 1), size(A_work, 2)) + cacheval.ipiv = similar( + A_work, BlasInt, min(size(A_work, 1), size(A_work, 2)) + ) + end + info_value = LinearSolve._direct_lu_factorize!(cacheval, A_work, alg) if info_value != 0 if verbose.blas_info != SciMLLogging.Silent() || verbose.blas_errors != SciMLLogging.Silent() || verbose.blas_invalid_args != SciMLLogging.Silent() - op_info = get_blas_operation_info(:dgetrf, A, cache.b, condition = verbose.condition_number != SciMLLogging.Silent()) - @SciMLMessage(cache.verbose, :condition_number) do - if op_info[:condition_number] === nothing - return "Matrix condition number calculation failed." - else - return "Matrix condition number: $(round(op_info[:condition_number], sigdigits = 4)) for $(size(A, 1))×$(size(A, 2)) matrix in dgetrf" + failure_op_info = get_blas_operation_info( + :dgetrf, A_work, cache.b, + condition = verbose.condition_number != SciMLLogging.Silent() + ) + let op_info = failure_op_info + @SciMLMessage(cache.verbose, :condition_number) do + if isinf(op_info.condition_number) + return "Matrix condition number calculation failed." + else + return "Matrix condition number: $(round(op_info.condition_number, sigdigits = 4)) for $(size(A_work, 1))×$(size(A_work, 2)) matrix in dgetrf" + end end end verb_option, message = blas_info_msg( - :dgetrf, info_value; extra_context = op_info + :dgetrf, info_value; extra_context = failure_op_info ) @SciMLMessage(message, verbose, verb_option) end else @SciMLMessage(cache.verbose, :blas_success) do - op_info = get_blas_operation_info( - :dgetrf, A, cache.b, + success_op_info = get_blas_operation_info( + :dgetrf, A_work, cache.b, condition = verbose.condition_number != SciMLLogging.Silent() ) - @SciMLMessage(cache.verbose, :condition_number) do - if op_info[:condition_number] === nothing - return "Matrix condition number calculation failed." - else - return "Matrix condition number: $(round(op_info[:condition_number], sigdigits = 4)) for $(size(A, 1))×$(size(A, 2)) matrix in dgetrf" + let op_info = success_op_info + @SciMLMessage(cache.verbose, :condition_number) do + if isinf(op_info.condition_number) + return "Matrix condition number calculation failed." + else + return "Matrix condition number: $(round(op_info.condition_number, sigdigits = 4)) for $(size(A_work, 1))×$(size(A_work, 2)) matrix in dgetrf" + end end end - return "BLAS LU factorization (dgetrf) completed successfully for $(op_info[:matrix_size]) matrix" + return "BLAS LU factorization (dgetrf) completed successfully for $(success_op_info.matrix_size) matrix" end end - if !LinearAlgebra.issuccess(fact[1]) + if info_value != 0 @SciMLMessage("Solver failed", cache.verbose, :solver_failure) return SciMLBase.build_linear_solution( alg, cache.u, nothing, nothing; retcode = ReturnCode.Failure @@ -330,16 +370,17 @@ function SciMLBase.solve!( cache.isfresh = false end - A, info = @get_cacheval(cache, :BLISLUFactorization) + cacheval = @get_cacheval(cache, :BLISLUFactorization) + factors = cacheval.factors + info = cacheval.info require_one_based_indexing(cache.u, cache.b) - m, n = size(A, 1), size(A, 2) + m, n = size(factors, 1), size(factors, 2) if m > n Bc = copy(cache.b) - getrs!('N', A.factors, A.ipiv, Bc; info) + getrs!('N', factors, cacheval.ipiv, Bc, info) copyto!(cache.u, 1, Bc, 1, n) else - copyto!(cache.u, cache.b) - getrs!('N', A.factors, A.ipiv, cache.u; info) + LinearSolve._direct_lu_solve!(cacheval, cache.u, cache.b, alg) end return SciMLBase.build_linear_solution(alg, cache.u, nothing, nothing; retcode = ReturnCode.Success) diff --git a/ext/LinearSolveChainRulesCoreExt.jl b/ext/LinearSolveChainRulesCoreExt.jl index ee0aa3278..87bd8d537 100644 --- a/ext/LinearSolveChainRulesCoreExt.jl +++ b/ext/LinearSolveChainRulesCoreExt.jl @@ -7,7 +7,7 @@ using LinearSolve: SciMLLinearSolveAlgorithm, AbstractFactorization, using SciMLBase: SciMLBase, LinearProblem, init, solve, solve! using SciMLOperators: issquare using ChainRulesCore: ChainRulesCore, NoTangent -using LinearAlgebra: Factorization, adjoint +using LinearAlgebra: adjoint const CRC = ChainRulesCore @@ -31,16 +31,20 @@ function CRC.rrule( @assert sensealg isa LinearSolveAdjoint "Currently only `LinearSolveAdjoint` is supported for adjoint sensitivity analysis." - # Decide if we need to cache `A` and `b` for the reverse pass + A_ = nothing if sensealg.linsolve === missing - # We can reuse the factorization so no copy is needed - # Krylov Methods don't modify `A`, so it's safe to just reuse it - # No Copy is needed even for the default case + can_reuse_factorization = LinearSolve._can_reuse_cache_factorization( + alg, cache.cacheval + ) if !( - alg isa AbstractFactorization || alg isa AbstractKrylovSubspaceMethod || + can_reuse_factorization || alg isa AbstractKrylovSubspaceMethod || alg isa DefaultLinearSolver ) - A_ = alias_A ? deepcopy(A) : A + A_ = if alg isa AbstractFactorization + deepcopy(A) + else + alias_A ? deepcopy(A) : A + end end else A_ = deepcopy(A) @@ -53,13 +57,15 @@ function CRC.rrule( ∂u = hasproperty(∂sol, :u) ? ∂sol.u : ∂sol if sensealg.linsolve === missing - λ = if cache.cacheval isa Factorization - cache.cacheval' \ ∂u - elseif cache.cacheval isa Tuple && cache.cacheval[1] isa Factorization - first(cache.cacheval)' \ ∂u + cached_adjoint_solution = LinearSolve._adjoint_factorization_solve( + alg, cache.cacheval, cache.A, ∂u + ) + λ = if cached_adjoint_solution !== nothing + cached_adjoint_solution elseif alg isa AbstractKrylovSubspaceMethod - invprob = LinearProblem(adjoint(cache.A), ∂u) - solve(invprob, alg; cache.abstol, cache.reltol, cache.verbose).u + LinearSolve._adjoint_krylov_solve( + alg, cache.A, ∂u; cache.abstol, cache.reltol, cache.verbose + ) elseif alg isa DefaultLinearSolver LinearSolve.defaultalg_adjoint_eval(cache, ∂u) else diff --git a/ext/LinearSolveEnzymeExt.jl b/ext/LinearSolveEnzymeExt.jl index 4ae616090..39220602f 100644 --- a/ext/LinearSolveEnzymeExt.jl +++ b/ext/LinearSolveEnzymeExt.jl @@ -716,15 +716,15 @@ function EnzymeRules.reverse( # Add the contribution from direct `linsolve.u` modifications dy .+= dy2.u - z = if _linsolve.cacheval isa Factorization - _linsolve.cacheval' \ dy - elseif _linsolve.cacheval isa Tuple && _linsolve.cacheval[1] isa Factorization - _linsolve.cacheval[1]' \ dy + cached_adjoint_solution = LinearSolve._adjoint_factorization_solve( + _linsolve.alg, _linsolve.cacheval, _linsolve.A, dy + ) + z = if cached_adjoint_solution !== nothing + cached_adjoint_solution elseif _linsolve.alg isa LinearSolve.AbstractKrylovSubspaceMethod # Doesn't modify `A`, so it's safe to just reuse it - invprob = LinearSolve.LinearProblem(transpose(_linsolve.A), dy) - solve( - invprob, _linsolve.alg; + LinearSolve._adjoint_krylov_solve( + _linsolve.alg, _linsolve.A, dy; abstol = _linsolve.abstol, reltol = _linsolve.reltol, verbose = _linsolve.verbose diff --git a/ext/LinearSolveFastLapackInterfaceExt.jl b/ext/LinearSolveFastLapackInterfaceExt.jl index 3c7fb9a2f..5d8172f65 100644 --- a/ext/LinearSolveFastLapackInterfaceExt.jl +++ b/ext/LinearSolveFastLapackInterfaceExt.jl @@ -9,6 +9,10 @@ struct WorkspaceAndFactors{W, F} factors::F end +LinearSolve._custom_cache_factorization( + ::Union{FastLUFactorization, FastQRFactorization}, cacheval::WorkspaceAndFactors +) = cacheval.factors + function LinearSolve.init_cacheval( ::FastLUFactorization, A, b, u, Pl, Pr, maxiters::Int, abstol, reltol, verbose::Union{LinearVerbosity, Bool}, diff --git a/ext/LinearSolveHSLExt.jl b/ext/LinearSolveHSLExt.jl index b04c8d32c..f15fe7bc9 100644 --- a/ext/LinearSolveHSLExt.jl +++ b/ext/LinearSolveHSLExt.jl @@ -138,6 +138,19 @@ function SciMLBase.solve!( end end +LinearSolve._custom_can_reuse_adjoint_factorization( + ::LinearSolve.HSLMA57Factorization, ::HSLMA57Cache +) = true + +function LinearSolve._custom_adjoint_factorization_solve( + ::LinearSolve.HSLMA57Factorization, hcache::HSLMA57Cache, A, b + ) + solution = copy(b) + _resize_ma57_work!(hcache, b) + HSL.ma57_solve!(hcache.ma57, solution, hcache.work) + return solution +end + function SciMLBase.solve!( cache::LinearSolve.LinearCache, alg::LinearSolve.HSLMA97Factorization; @@ -187,4 +200,16 @@ function SciMLBase.solve!( end end +LinearSolve._custom_can_reuse_adjoint_factorization( + ::LinearSolve.HSLMA97Factorization, ::HSLMA97Cache +) = true + +function LinearSolve._custom_adjoint_factorization_solve( + ::LinearSolve.HSLMA97Factorization, hcache::HSLMA97Cache, A, b + ) + solution = copy(b) + HSL.ma97_solve!(hcache.ma97, solution) + return solution +end + end diff --git a/ext/LinearSolveMUMPSExt.jl b/ext/LinearSolveMUMPSExt.jl index 1590a1da0..6ec4fd268 100644 --- a/ext/LinearSolveMUMPSExt.jl +++ b/ext/LinearSolveMUMPSExt.jl @@ -173,4 +173,29 @@ function SciMLBase.solve!( ) end +LinearSolve._custom_can_reuse_adjoint_factorization( + ::LinearSolve.MUMPSFactorization, cache::MUMPSCache +) = true + +function LinearSolve._custom_adjoint_factorization_solve( + alg::LinearSolve.MUMPSFactorization, cache::MUMPSCache, A, b + ) + solver = cache.solver + solver === nothing && return nothing + solution = similar(b) + reverse_transposed = !alg.transposed + # MUMPS exposes transpose but not adjoint solves. Conjugating both sides + # turns a transpose solve with the cached factors into an adjoint solve. + rhs = eltype(A) <: Real ? b : conj.(b) + MUMPS.associate_rhs!(solver, rhs) + reverse_transposed && transpose!(solver) + try + MUMPS.mumps_solve!(solver; rhs_changed = true) + MUMPS.get_sol!(solution, solver) + finally + reverse_transposed && transpose!(solver) + end + return eltype(A) <: Real ? solution : conj.(solution) +end + end diff --git a/ext/LinearSolveMooncakeExt.jl b/ext/LinearSolveMooncakeExt.jl index 6676b4284..d02f21730 100644 --- a/ext/LinearSolveMooncakeExt.jl +++ b/ext/LinearSolveMooncakeExt.jl @@ -69,13 +69,17 @@ function Mooncake.rrule!!( @assert sensealg isa LinearSolveAdjoint "Currently only `LinearSolveAdjoint` is supported for adjoint sensitivity analysis." - # logic behind caching `A` and `b` for the reverse pass based on rrule above for SciMLBase.solve + A_ = nothing if sensealg.linsolve === missing + can_reuse_factorization = LinearSolve._can_reuse_cache_factorization( + alg, cache.cacheval + ) if !( - alg isa LinearSolve.AbstractFactorization || alg isa LinearSolve.AbstractKrylovSubspaceMethod || + can_reuse_factorization || alg isa LinearSolve.AbstractKrylovSubspaceMethod || alg isa LinearSolve.DefaultLinearSolver ) - A_ = alias_A ? deepcopy(A) : A + A_ = alg isa LinearSolve.AbstractFactorization ? deepcopy(A) : + alias_A ? deepcopy(A) : A end else A_ = deepcopy(A) @@ -93,13 +97,15 @@ function Mooncake.rrule!!( ∂u = sol.dx.data.u if sensealg.linsolve === missing - λ = if cache.cacheval isa Factorization - cache.cacheval' \ ∂u - elseif cache.cacheval isa Tuple && cache.cacheval[1] isa Factorization - first(cache.cacheval)' \ ∂u + cached_adjoint_solution = LinearSolve._adjoint_factorization_solve( + alg, cache.cacheval, cache.A, ∂u + ) + λ = if cached_adjoint_solution !== nothing + cached_adjoint_solution elseif alg isa AbstractKrylovSubspaceMethod - invprob = LinearProblem(adjoint(cache.A), ∂u) - solve(invprob, alg; cache.abstol, cache.reltol, cache.verbose).u + LinearSolve._adjoint_krylov_solve( + alg, cache.A, ∂u; cache.abstol, cache.reltol, cache.verbose + ) elseif alg isa DefaultLinearSolver LinearSolve.defaultalg_adjoint_eval(cache, ∂u) else diff --git a/ext/LinearSolvePardisoExt.jl b/ext/LinearSolvePardisoExt.jl index 37248723c..4d561f726 100644 --- a/ext/LinearSolvePardisoExt.jl +++ b/ext/LinearSolvePardisoExt.jl @@ -161,6 +161,31 @@ function SciMLBase.solve!(cache::LinearSolve.LinearCache, alg::PardisoJL; kwargs return SciMLBase.build_linear_solution(alg, cache.u, nothing, nothing) end +LinearSolve._custom_can_reuse_adjoint_factorization( + ::PardisoJL, ::Pardiso.AbstractPardisoSolver +) = true + +function LinearSolve._custom_adjoint_factorization_solve( + ::PardisoJL, solver::Pardiso.AbstractPardisoSolver, A, b + ) + transposed_iparm = Pardiso.get_iparm(solver, 12) + solution = similar(b) + # Pardiso sees the CSC storage as CSR for transpose(A). With transpose mode + # disabled, conjugating both sides solves the adjoint system for complex A. + rhs = eltype(A) <: Real ? b : conj.(b) + Pardiso.set_iparm!(solver, 12, 0) + Pardiso.set_phase!(solver, Pardiso.SOLVE_ITERATIVE_REFINE) + try + Pardiso.pardiso( + solver, solution, + SparseMatrixCSC(size(A)..., getcolptr(A), rowvals(A), nonzeros(A)), rhs + ) + finally + Pardiso.set_iparm!(solver, 12, transposed_iparm) + end + return eltype(A) <: Real ? solution : conj.(solution) +end + # Add finalizer to release memory # Pardiso.set_phase!(cache.cacheval, Pardiso.RELEASE_ALL) diff --git a/ext/LinearSolveRecursiveFactorizationExt.jl b/ext/LinearSolveRecursiveFactorizationExt.jl index 3de6dbf56..e387be470 100644 --- a/ext/LinearSolveRecursiveFactorizationExt.jl +++ b/ext/LinearSolveRecursiveFactorizationExt.jl @@ -157,4 +157,23 @@ function LinearSolve.init_cacheval( return ws = RecursiveFactorization.🦋workspace(A, b), RecursiveFactorization.lu!(rand(1, 1), Val(false), alg.thread) end +LinearSolve._custom_can_reuse_adjoint_factorization( + ::ButterflyFactorization, ::Tuple +) = true + +function LinearSolve._custom_adjoint_factorization_solve( + ::ButterflyFactorization, cacheval::Tuple, A, b::AbstractVector + ) + workspace, factorization = cacheval + n = workspace.n + T = promote_type(eltype(workspace.U), eltype(b)) + padded_rhs = zeros(T, size(workspace.U, 1)) + copyto!(view(padded_rhs, 1:n), b) + transformed_rhs = adjoint(workspace.V) * padded_rhs + upper_solution = adjoint(factorization.U) \ transformed_rhs + factorization_solution = adjoint(factorization.L) \ upper_solution + solution = workspace.U * factorization_solution + return solution[1:n] +end + end diff --git a/ext/LinearSolveSparseArraysExt.jl b/ext/LinearSolveSparseArraysExt.jl index c74018852..bd8951b64 100644 --- a/ext/LinearSolveSparseArraysExt.jl +++ b/ext/LinearSolveSparseArraysExt.jl @@ -779,6 +779,19 @@ function SciMLBase.solve!( ) end +LinearSolve._custom_can_reuse_adjoint_factorization( + ::SparseColumnPivotedQRFactorization, + ::SCPQR.SparseColumnPivotedQRFactorization +) = true + +function LinearSolve._custom_adjoint_factorization_solve( + ::SparseColumnPivotedQRFactorization, + factorization::SCPQR.SparseColumnPivotedQRFactorization, + A, b + ) + return adjoint(factorization) \ b +end + # SparseColumnPivotedQR's ldiv! only accepts vector right-hand sides; batched # (matrix) right-hand sides solve column-by-column against the one factorization. function LinearSolve._ldiv!( diff --git a/src/LinearSolve.jl b/src/LinearSolve.jl index a402e159a..65dd4a896 100644 --- a/src/LinearSolve.jl +++ b/src/LinearSolve.jl @@ -418,6 +418,7 @@ include("appleaccelerate.jl") include("mkl.jl") include("openblas.jl") include("simplelu.jl") +include("adjoint_factorization.jl") include("simplegmres.jl") include("iterative_wrappers.jl") include("preconditioners.jl") diff --git a/src/adjoint_factorization.jl b/src/adjoint_factorization.jl new file mode 100644 index 000000000..08227f6e2 --- /dev/null +++ b/src/adjoint_factorization.jl @@ -0,0 +1,224 @@ +abstract type _AdjointFactorizationReuse end +struct _DirectAdjointFactorizationReuse <: _AdjointFactorizationReuse end +struct _ExtractedAdjointFactorizationReuse <: _AdjointFactorizationReuse end +struct _NormalAdjointFactorizationReuse <: _AdjointFactorizationReuse end +struct _CustomAdjointFactorizationReuse <: _AdjointFactorizationReuse end +struct _NoAdjointFactorizationReuse <: _AdjointFactorizationReuse end +struct _UnspecifiedAdjointFactorizationReuse <: _AdjointFactorizationReuse end + +_adjoint_factorization_reuse(::Type{<:SciMLLinearSolveAlgorithm}) = + _NoAdjointFactorizationReuse() +_adjoint_factorization_reuse(::Type{<:AbstractFactorization}) = + _UnspecifiedAdjointFactorizationReuse() + +for Alg in ( + LUFactorization, + GenericLUFactorization, + GESVFactorization, + QRFactorization, + CholeskyFactorization, + LDLtFactorization, + SVDFactorization, + BunchKaufmanFactorization, + GenericFactorization, + UMFPACKFactorization, + KLUFactorization, + PureKLUFactorization, + CHOLMODFactorization, + CliqueTreesFactorization, + RFLUFactorization, + MKLLUFactorization, + MetalLUFactorization, + MetalOffload32MixedLUFactorization, + MKL32MixedLUFactorization, + AppleAccelerate32MixedLUFactorization, + OpenBLAS32MixedLUFactorization, + RF32MixedLUFactorization, + ) + @eval _adjoint_factorization_reuse(::Type{<:$Alg}) = + _DirectAdjointFactorizationReuse() +end + +for Alg in ( + AppleAccelerateLUFactorization, + OpenBLASLUFactorization, + BLISLUFactorization, + FastLUFactorization, + FastQRFactorization, + ) + @eval _adjoint_factorization_reuse(::Type{<:$Alg}) = + _ExtractedAdjointFactorizationReuse() +end + +for Alg in (NormalCholeskyFactorization, NormalBunchKaufmanFactorization) + @eval _adjoint_factorization_reuse(::Type{<:$Alg}) = + _NormalAdjointFactorizationReuse() +end + +for Alg in ( + SimpleLUFactorization, + SparseColumnPivotedQRFactorization, + ButterflyFactorization, + PardisoJL, + MUMPSFactorization, + HSLMA57Factorization, + HSLMA97Factorization, + ) + @eval _adjoint_factorization_reuse(::Type{<:$Alg}) = + _CustomAdjointFactorizationReuse() +end + +# These integrations either do not cache a numeric factorization or their public +# backend API does not expose an adjoint solve using the cached factorization. +# Generic reverse paths preserve a copy of `A` and factorize `adjoint(A)`; +# AD backends without that fallback report the algorithm as unsupported. +for Alg in ( + DiagonalFactorization, + PureUMFPACKFactorization, + SparspakFactorization, + STRUMPACKFactorization, + CudaOffloadLUFactorization, + CUDAOffload32MixedLUFactorization, + CudaOffloadQRFactorization, + CudaOffloadFactorization, + AMDGPUOffloadLUFactorization, + AMDGPUOffloadQRFactorization, + CUSOLVERRFFactorization, + ParUFactorization, + SuperLUDISTFactorization, + ElementalJL, + SpecializedLUFactorization, + SpecializedQRFactorization, + ) + @eval _adjoint_factorization_reuse(::Type{<:$Alg}) = + _NoAdjointFactorizationReuse() +end + +function _standard_cache_factorization(cacheval) + if cacheval isa Factorization + return cacheval + elseif cacheval isa Tuple && !isempty(cacheval) && first(cacheval) isa Factorization + return first(cacheval) + else + return nothing + end +end + +_custom_cache_factorization(::AbstractFactorization, cacheval) = nothing + +""" + _cache_factorization(alg, cacheval) + +Return the factorization exposed by an explicitly opted-in algorithm, or +`nothing`. The direct-cache fallback accepts a `LinearAlgebra.Factorization` +stored either directly or first in a tuple, but it is only used for algorithms +listed as `_DirectAdjointFactorizationReuse`. Custom cache layouts opt in with +`_custom_cache_factorization`. +""" +function _cache_factorization(alg::AbstractFactorization, cacheval) + reuse = _adjoint_factorization_reuse(typeof(alg)) + return _cache_factorization(reuse, alg, cacheval) +end +_cache_factorization(::SciMLLinearSolveAlgorithm, cacheval) = nothing +_cache_factorization(::_DirectAdjointFactorizationReuse, alg, cacheval) = + _standard_cache_factorization(cacheval) +_cache_factorization(::_ExtractedAdjointFactorizationReuse, alg, cacheval) = + _custom_cache_factorization(alg, cacheval) +_cache_factorization(::_AdjointFactorizationReuse, alg, cacheval) = nothing + +function _can_reuse_cache_factorization(alg::AbstractFactorization, cacheval) + reuse = _adjoint_factorization_reuse(typeof(alg)) + return _can_reuse_cache_factorization(reuse, alg, cacheval) +end +_can_reuse_cache_factorization(::SciMLLinearSolveAlgorithm, cacheval) = false +_can_reuse_cache_factorization(::_DirectAdjointFactorizationReuse, alg, cacheval) = + _standard_cache_factorization(cacheval) !== nothing +_can_reuse_cache_factorization(::_ExtractedAdjointFactorizationReuse, alg, cacheval) = + _custom_cache_factorization(alg, cacheval) !== nothing +_can_reuse_cache_factorization(::_NormalAdjointFactorizationReuse, alg, cacheval) = + _standard_cache_factorization(cacheval) !== nothing +_can_reuse_cache_factorization(::_CustomAdjointFactorizationReuse, alg, cacheval) = + _custom_can_reuse_adjoint_factorization(alg, cacheval) +_can_reuse_cache_factorization(::_AdjointFactorizationReuse, alg, cacheval) = false + +_custom_can_reuse_adjoint_factorization(::AbstractFactorization, cacheval) = false +_custom_adjoint_factorization_solve(::AbstractFactorization, cacheval, A, b) = nothing + +""" + _adjoint_factorization_solve(alg, cacheval, A, b) + +Solve `adjoint(A) * x = b` using `alg`'s cached factorization, returning +`nothing` when that algorithm has not opted into reverse-pass reuse. Algorithms +whose cache does not directly represent `A` provide a solver-specific method. +""" +function _adjoint_factorization_solve( + alg::AbstractFactorization, cacheval, A, b + ) + reuse = _adjoint_factorization_reuse(typeof(alg)) + return _adjoint_factorization_solve(reuse, alg, cacheval, A, b) +end +_adjoint_factorization_solve(::SciMLLinearSolveAlgorithm, cacheval, A, b) = nothing + +function _adjoint_factorization_solve( + ::Union{_DirectAdjointFactorizationReuse, _ExtractedAdjointFactorizationReuse}, + alg, cacheval, A, b + ) + factorization = _cache_factorization(alg, cacheval) + return factorization === nothing ? nothing : factorization' \ b +end + +function _adjoint_factorization_solve( + ::_NormalAdjointFactorizationReuse, alg, cacheval, A, b + ) + factorization = _standard_cache_factorization(cacheval) + return factorization === nothing ? nothing : A * (factorization \ b) +end + +function _adjoint_factorization_solve( + ::_CustomAdjointFactorizationReuse, alg, cacheval, A, b + ) + return _custom_adjoint_factorization_solve(alg, cacheval, A, b) +end + +_adjoint_factorization_solve(::_AdjointFactorizationReuse, alg, cacheval, A, b) = + nothing + +function _adjoint_krylov_solve( + alg::AbstractKrylovSubspaceMethod, A, b; abstol, reltol, verbose + ) + invprob = LinearProblem(adjoint(A), b) + return solve(invprob, alg; abstol, reltol, verbose).u +end + +_custom_can_reuse_adjoint_factorization(::SimpleLUFactorization, ::LUSolver) = true + +function _custom_adjoint_factorization_solve( + ::SimpleLUFactorization, factorization::LUSolver, A, b::AbstractVector + ) + n = factorization.n + T = promote_type(eltype(factorization.A), eltype(b)) + y = similar(b, T, n) + z = similar(b, T, n) + x = similar(b, T, n) + + @inbounds for i in 1:n + value = b[i] + for j in 1:(i - 1) + value -= conj(factorization.A[j, i]) * y[j] + end + y[i] = value / conj(factorization.A[i, i]) + end + + @inbounds for i in n:-1:1 + value = y[i] + for j in (i + 1):n + value -= conj(factorization.A[j, i]) * z[j] + end + z[i] = value + end + + @inbounds for i in 1:n + x[factorization.perms[i]] = z[i] + end + return x +end diff --git a/src/appleaccelerate.jl b/src/appleaccelerate.jl index 72bc773c4..73e24943e 100644 --- a/src/appleaccelerate.jl +++ b/src/appleaccelerate.jl @@ -31,11 +31,9 @@ else end end -function aa_getrf!( - A::AbstractMatrix{<:ComplexF64}; - ipiv = similar(A, Cint, min(size(A, 1), size(A, 2))), - info = Ref{Cint}(), - check = false +@inline function aa_getrf!( + A::AbstractMatrix{<:ComplexF64}, ipiv::AbstractVector{Cint}, + info::Ref{Cint}, check::Bool ) __appleaccelerate_isavailable() || error("Error, AppleAccelerate binary is missing but solve is being called. Report this issue") @@ -44,9 +42,7 @@ function aa_getrf!( chkstride1(A) m, n = size(A) lda = max(1, stride(A, 2)) - if isempty(ipiv) - ipiv = similar(A, Cint, min(size(A, 1), size(A, 2))) - end + length(ipiv) == min(m, n) || throw(DimensionMismatch("ipiv has length $(length(ipiv)), but needs $(min(m, n))")) ccall( ("zgetrf_", libacc), Cvoid, ( @@ -56,14 +52,12 @@ function aa_getrf!( m, n, A, lda, ipiv, info ) info[] < 0 && throw(ArgumentError("Invalid arguments sent to LAPACK dgetrf_")) - return A, ipiv, BlasInt(info[]), info #Error code is stored in LU factorization type + return BlasInt(info[]) end -function aa_getrf!( - A::AbstractMatrix{<:ComplexF32}; - ipiv = similar(A, Cint, min(size(A, 1), size(A, 2))), - info = Ref{Cint}(), - check = false +@inline function aa_getrf!( + A::AbstractMatrix{<:ComplexF32}, ipiv::AbstractVector{Cint}, + info::Ref{Cint}, check::Bool ) __appleaccelerate_isavailable() || error("Error, AppleAccelerate binary is missing but solve is being called. Report this issue") @@ -72,9 +66,7 @@ function aa_getrf!( chkstride1(A) m, n = size(A) lda = max(1, stride(A, 2)) - if isempty(ipiv) - ipiv = similar(A, Cint, min(size(A, 1), size(A, 2))) - end + length(ipiv) == min(m, n) || throw(DimensionMismatch("ipiv has length $(length(ipiv)), but needs $(min(m, n))")) ccall( ("cgetrf_", libacc), Cvoid, ( @@ -84,14 +76,12 @@ function aa_getrf!( m, n, A, lda, ipiv, info ) info[] < 0 && throw(ArgumentError("Invalid arguments sent to LAPACK dgetrf_")) - return A, ipiv, BlasInt(info[]), info #Error code is stored in LU factorization type + return BlasInt(info[]) end -function aa_getrf!( - A::AbstractMatrix{<:Float64}; - ipiv = similar(A, Cint, min(size(A, 1), size(A, 2))), - info = Ref{Cint}(), - check = false +@inline function aa_getrf!( + A::AbstractMatrix{<:Float64}, ipiv::AbstractVector{Cint}, + info::Ref{Cint}, check::Bool ) __appleaccelerate_isavailable() || error("Error, AppleAccelerate binary is missing but solve is being called. Report this issue") @@ -100,9 +90,7 @@ function aa_getrf!( chkstride1(A) m, n = size(A) lda = max(1, stride(A, 2)) - if isempty(ipiv) - ipiv = similar(A, Cint, min(size(A, 1), size(A, 2))) - end + length(ipiv) == min(m, n) || throw(DimensionMismatch("ipiv has length $(length(ipiv)), but needs $(min(m, n))")) ccall( ("dgetrf_", libacc), Cvoid, ( @@ -112,14 +100,12 @@ function aa_getrf!( m, n, A, lda, ipiv, info ) info[] < 0 && throw(ArgumentError("Invalid arguments sent to LAPACK dgetrf_")) - return A, ipiv, BlasInt(info[]), info #Error code is stored in LU factorization type + return BlasInt(info[]) end -function aa_getrf!( - A::AbstractMatrix{<:Float32}; - ipiv = similar(A, Cint, min(size(A, 1), size(A, 2))), - info = Ref{Cint}(), - check = false +@inline function aa_getrf!( + A::AbstractMatrix{<:Float32}, ipiv::AbstractVector{Cint}, + info::Ref{Cint}, check::Bool ) __appleaccelerate_isavailable() || error("Error, AppleAccelerate binary is missing but solve is being called. Report this issue") @@ -128,9 +114,7 @@ function aa_getrf!( chkstride1(A) m, n = size(A) lda = max(1, stride(A, 2)) - if isempty(ipiv) - ipiv = similar(A, Cint, min(size(A, 1), size(A, 2))) - end + length(ipiv) == min(m, n) || throw(DimensionMismatch("ipiv has length $(length(ipiv)), but needs $(min(m, n))")) ccall( ("sgetrf_", libacc), Cvoid, @@ -141,15 +125,15 @@ function aa_getrf!( m, n, A, lda, ipiv, info ) info[] < 0 && throw(ArgumentError("Invalid arguments sent to LAPACK dgetrf_")) - return A, ipiv, BlasInt(info[]), info #Error code is stored in LU factorization type + return BlasInt(info[]) end function aa_getrs!( trans::AbstractChar, A::AbstractMatrix{<:ComplexF64}, ipiv::AbstractVector{Cint}, - B::AbstractVecOrMat{<:ComplexF64}; - info = Ref{Cint}() + B::AbstractVecOrMat{<:ComplexF64}, + info::Ref{Cint} ) __appleaccelerate_isavailable() || error("Error, AppleAccelerate binary is missing but solve is being called. Report this issue") @@ -180,8 +164,8 @@ function aa_getrs!( trans::AbstractChar, A::AbstractMatrix{<:ComplexF32}, ipiv::AbstractVector{Cint}, - B::AbstractVecOrMat{<:ComplexF32}; - info = Ref{Cint}() + B::AbstractVecOrMat{<:ComplexF32}, + info::Ref{Cint} ) __appleaccelerate_isavailable() || error("Error, AppleAccelerate binary is missing but solve is being called. Report this issue") @@ -213,8 +197,8 @@ function aa_getrs!( trans::AbstractChar, A::AbstractMatrix{<:Float64}, ipiv::AbstractVector{Cint}, - B::AbstractVecOrMat{<:Float64}; - info = Ref{Cint}() + B::AbstractVecOrMat{<:Float64}, + info::Ref{Cint} ) __appleaccelerate_isavailable() || error("Error, AppleAccelerate binary is missing but solve is being called. Report this issue") @@ -246,8 +230,8 @@ function aa_getrs!( trans::AbstractChar, A::AbstractMatrix{<:Float32}, ipiv::AbstractVector{Cint}, - B::AbstractVecOrMat{<:Float32}; - info = Ref{Cint}() + B::AbstractVecOrMat{<:Float32}, + info::Ref{Cint} ) __appleaccelerate_isavailable() || error("Error, AppleAccelerate binary is missing but solve is being called. Report this issue") @@ -280,10 +264,30 @@ _get_residualsafety(alg::AppleAccelerateLUFactorization) = alg.residualsafety default_alias_A(::AppleAccelerateLUFactorization, ::Any, ::Any) = false default_alias_b(::AppleAccelerateLUFactorization, ::Any, ::Any) = false -const PREALLOCATED_APPLE_LU = begin - A = rand(0, 0) - luinst = ArrayInterface.lu_instance(A) - LU(luinst.factors, similar(A, Cint, 0), luinst.info), Ref{Cint}() +mutable struct AppleAccelerateLUCache{F, P, I} + factors::F + ipiv::P + info::I +end + +_custom_cache_factorization( + ::AppleAccelerateLUFactorization, cacheval::AppleAccelerateLUCache +) = + LU(cacheval.factors, BlasInt.(cacheval.ipiv), Int(cacheval.info[])) + +@inline function _direct_lu_factorize!( + cacheval::AppleAccelerateLUCache, A_work, ::AppleAccelerateLUFactorization + ) + cacheval.factors = A_work + return aa_getrf!(A_work, cacheval.ipiv, cacheval.info, false) +end + +@inline function _direct_lu_solve!( + cacheval::AppleAccelerateLUCache, u, b, ::AppleAccelerateLUFactorization + ) + copyto!(u, b) + aa_getrs!('N', cacheval.factors, cacheval.ipiv, u, cacheval.info) + return u end function LinearSolve.init_cacheval( @@ -291,7 +295,8 @@ function LinearSolve.init_cacheval( maxiters::Int, abstol, reltol, verbose::Union{LinearVerbosity, Bool}, assumptions::OperatorAssumptions ) - return PREALLOCATED_APPLE_LU + A0 = Matrix{Float64}(undef, 0, 0) + return AppleAccelerateLUCache(A0, Vector{Cint}(undef, 0), Ref{Cint}()) end function LinearSolve.init_cacheval( @@ -300,10 +305,9 @@ function LinearSolve.init_cacheval( maxiters::Int, abstol, reltol, verbose::Union{LinearVerbosity, Bool}, assumptions::OperatorAssumptions ) - # Ask `lu_instance` about `A` itself so wrapper-specific dispatch can choose - # the factorization container that `solve!` will store. - luinst = ArrayInterface.lu_instance(A) - return LU(luinst.factors, similar(luinst.ipiv, Cint, 0), luinst.info), Ref{Cint}() + return AppleAccelerateLUCache( + A, similar(A, Cint, min(size(A, 1), size(A, 2))), Ref{Cint}() + ) end function SciMLBase.solve!( @@ -312,60 +316,64 @@ function SciMLBase.solve!( ) __appleaccelerate_isavailable() || error("Error, AppleAccelerate binary is missing but solve is being called. Report this issue") - A = cache.A - A = convert(AbstractMatrix, A) + A_work = convert(AbstractMatrix, cache.A) check_safety = alg.residualsafety && cache.isfresh needs_backup = check_safety || (cache.alg isa DefaultLinearSolver && cache.alg.safetyfallback && cache.isfresh) - A_original = needs_backup ? _copy_A_for_safety(cache) : A + A_original = needs_backup ? _copy_A_for_safety(cache) : A_work verbose = cache.verbose if cache.isfresh cacheval = @get_cacheval(cache, :AppleAccelerateLUFactorization) - res = aa_getrf!(A; ipiv = cacheval[1].ipiv, info = cacheval[2]) - fact = LU(res[1:3]...), res[4] - cache.cacheval = fact - - info_value = res[3] + if length(cacheval.ipiv) != min(size(A_work, 1), size(A_work, 2)) + cacheval.ipiv = similar( + A_work, Cint, min(size(A_work, 1), size(A_work, 2)) + ) + end + info_value = _direct_lu_factorize!(cacheval, A_work, alg) if info_value != 0 if verbose.blas_info != SciMLLogging.Silent() || verbose.blas_errors != SciMLLogging.Silent() || verbose.blas_invalid_args != SciMLLogging.Silent() - op_info = get_blas_operation_info( - :dgetrf, A, cache.b, + failure_op_info = get_blas_operation_info( + :dgetrf, A_work, cache.b, condition = verbose.condition_number != SciMLLogging.Silent() ) - @SciMLMessage(cache.verbose, :condition_number) do - if isinf(op_info.condition_number) - return "Matrix condition number calculation failed." - else - return "Matrix condition number: $(round(op_info.condition_number, sigdigits = 4)) for $(size(A, 1))×$(size(A, 2)) matrix in dgetrf" + let op_info = failure_op_info + @SciMLMessage(cache.verbose, :condition_number) do + if isinf(op_info.condition_number) + return "Matrix condition number calculation failed." + else + return "Matrix condition number: $(round(op_info.condition_number, sigdigits = 4)) for $(size(A_work, 1))×$(size(A_work, 2)) matrix in dgetrf" + end end end verb_option, message = blas_info_msg( - :dgetrf, info_value; extra_context = op_info + :dgetrf, info_value; extra_context = failure_op_info ) @SciMLMessage(message, verbose, verb_option) end else @SciMLMessage(cache.verbose, :blas_success) do - op_info = get_blas_operation_info( - :dgetrf, A, cache.b, + success_op_info = get_blas_operation_info( + :dgetrf, A_work, cache.b, condition = verbose.condition_number != SciMLLogging.Silent() ) - @SciMLMessage(cache.verbose, :condition_number) do - if isinf(op_info.condition_number) - return "Matrix condition number calculation failed." - else - return "Matrix condition number: $(round(op_info.condition_number, sigdigits = 4)) for $(size(A, 1))×$(size(A, 2)) matrix in dgetrf" + let op_info = success_op_info + @SciMLMessage(cache.verbose, :condition_number) do + if isinf(op_info.condition_number) + return "Matrix condition number calculation failed." + else + return "Matrix condition number: $(round(op_info.condition_number, sigdigits = 4)) for $(size(A_work, 1))×$(size(A_work, 2)) matrix in dgetrf" + end end end - return "BLAS LU factorization (dgetrf) completed successfully for $(op_info.matrix_size) matrix" + return "BLAS LU factorization (dgetrf) completed successfully for $(success_op_info.matrix_size) matrix" end end - if !LinearAlgebra.issuccess(fact[1]) + if info_value != 0 @SciMLMessage("Solver failed", cache.verbose, :solver_failure) return SciMLBase.build_linear_solution( alg, cache.u, nothing, nothing; retcode = ReturnCode.Failure @@ -374,20 +382,21 @@ function SciMLBase.solve!( cache.isfresh = false end - A, info = @get_cacheval(cache, :AppleAccelerateLUFactorization) + cacheval = @get_cacheval(cache, :AppleAccelerateLUFactorization) + factors = cacheval.factors + info = cacheval.info require_one_based_indexing(cache.u, cache.b) - m, n = size(A, 1), size(A, 2) + m, n = size(factors, 1), size(factors, 2) if m > n Bc = copy(cache.b) - aa_getrs!('N', A.factors, A.ipiv, Bc; info) + aa_getrs!('N', factors, cacheval.ipiv, Bc, info) if cache.b isa AbstractMatrix copyto!(cache.u, @view(Bc[1:n, :])) else copyto!(cache.u, 1, Bc, 1, n) end else - copyto!(cache.u, cache.b) - aa_getrs!('N', A.factors, A.ipiv, cache.u; info) + _direct_lu_solve!(cacheval, cache.u, cache.b, alg) end if check_safety @@ -421,9 +430,10 @@ function LinearSolve.init_cacheval( A_32 = similar(A, T32) b_32 = similar(b, T32) u_32 = similar(u, T32) - luinst = ArrayInterface.lu_instance(rand(T32, 0, 0)) + ipiv = similar(A_32, Cint, min(size(A_32, 1), size(A_32, 2))) + luinst = LU(A_32, ipiv, zero(BlasInt)) # Return tuple with pre-allocated arrays - return (LU(luinst.factors, similar(A_32, Cint, 0), luinst.info), Ref{Cint}(), A_32, b_32, u_32) + return (luinst, Ref{Cint}(), A_32, b_32, u_32) end function SciMLBase.solve!( @@ -442,11 +452,9 @@ function SciMLBase.solve!( # Compute 32-bit type on demand and copy A T32 = eltype(A) <: Complex ? ComplexF32 : Float32 A_32 .= T32.(A) - res = aa_getrf!(A_32; ipiv = luinst.ipiv, info = info) - fact = (LU(res[1:3]...), res[4], A_32, b_32, u_32) - cache.cacheval = fact + info_value = aa_getrf!(A_32, luinst.ipiv, info, false) - if !LinearAlgebra.issuccess(fact[1]) + if info_value != 0 @SciMLMessage("Solver failed", cache.verbose, :solver_failure) return SciMLBase.build_linear_solution( alg, cache.u, nothing, nothing; retcode = ReturnCode.Failure @@ -468,7 +476,7 @@ function SciMLBase.solve!( b_32 .= T32.(cache.b) if m > n - aa_getrs!('N', A_lu.factors, A_lu.ipiv, b_32; info) + aa_getrs!('N', A_lu.factors, A_lu.ipiv, b_32, info) # Convert back to original precision if cache.b isa AbstractMatrix cache.u .= Torig.(@view(b_32[1:n, :])) @@ -477,7 +485,7 @@ function SciMLBase.solve!( end else copyto!(u_32, b_32) - aa_getrs!('N', A_lu.factors, A_lu.ipiv, u_32; info) + aa_getrs!('N', A_lu.factors, A_lu.ipiv, u_32, info) # Convert back to original precision cache.u .= Torig.(u_32) end diff --git a/src/default.jl b/src/default.jl index 06a695168..1ecbcc0d5 100644 --- a/src/default.jl +++ b/src/default.jl @@ -1117,8 +1117,8 @@ end end elseif alg == Symbol(DefaultAlgorithmChoice.AppleAccelerateLUFactorization) quote - A = getproperty(cache.cacheval, $(Meta.quot(alg)))[1] - aa_getrs!('T', A.factors, A.ipiv, dy) + A = getproperty(cache.cacheval, $(Meta.quot(alg))) + aa_getrs!('T', A.factors, A.ipiv, dy, A.info) end elseif alg in Symbol.( ( diff --git a/src/factorization.jl b/src/factorization.jl index feefb0dfb..f0d586563 100644 --- a/src/factorization.jl +++ b/src/factorization.jl @@ -152,6 +152,9 @@ _ldiv!(x, A, b::SVector) = (x .= A \ b) _ldiv!(::SVector, A, b::SVector) = (A \ b) _ldiv!(::SVector, A, b) = (A \ b) +function _direct_lu_factorize! end +function _direct_lu_solve! end + # Build a column-pivoted sparse QR factorization of `A` (the default sparse-LU # singular fallback). The method is provided by the SparseArrays extension over # SparseColumnPivotedQR.jl; this generic declaration lets `src/default.jl` call it. diff --git a/src/init.jl b/src/init.jl index f01637a24..64fc9acf7 100644 --- a/src/init.jl +++ b/src/init.jl @@ -1,4 +1,5 @@ function __init__() + _init_openblas_symbols!() IS_OPENBLAS[] = occursin("openblas", BLAS.get_config().loaded_libs[1].libname) return HAS_APPLE_ACCELERATE[] = __appleaccelerate_isavailable() diff --git a/src/openblas.jl b/src/openblas.jl index f1a7db220..1c7294c2f 100644 --- a/src/openblas.jl +++ b/src/openblas.jl @@ -37,18 +37,45 @@ end OpenBLASLUFactorization(; residualsafety::Bool = false) = OpenBLASLUFactorization(residualsafety) -# Check if OpenBLAS is available -@static if !@isdefined(OpenBLAS_jll) - __openblas_isavailable() = false -else - __openblas_isavailable() = OpenBLAS_jll.is_available() +__openblas_isavailable() = useopenblas + +# Resolve Julia 1.13 lazy JLL products once so solves call fixed function pointers. +const _openblas_handle = Ref{Ptr{Cvoid}}(C_NULL) +const _openblas_zgetrf = Ref{Ptr{Cvoid}}(C_NULL) +const _openblas_cgetrf = Ref{Ptr{Cvoid}}(C_NULL) +const _openblas_dgetrf = Ref{Ptr{Cvoid}}(C_NULL) +const _openblas_sgetrf = Ref{Ptr{Cvoid}}(C_NULL) +const _openblas_zgetrs = Ref{Ptr{Cvoid}}(C_NULL) +const _openblas_cgetrs = Ref{Ptr{Cvoid}}(C_NULL) +const _openblas_dgetrs = Ref{Ptr{Cvoid}}(C_NULL) +const _openblas_sgetrs = Ref{Ptr{Cvoid}}(C_NULL) + +function _init_openblas_symbols!() + @static if VERSION >= v"1.13.0-DEV.0" && useopenblas + handle = Libdl.dlopen(libopenblas) + _openblas_handle[] = handle + _openblas_zgetrf[] = Libdl.dlsym(handle, @blasfunc(zgetrf_)) + _openblas_cgetrf[] = Libdl.dlsym(handle, @blasfunc(cgetrf_)) + _openblas_dgetrf[] = Libdl.dlsym(handle, @blasfunc(dgetrf_)) + _openblas_sgetrf[] = Libdl.dlsym(handle, @blasfunc(sgetrf_)) + _openblas_zgetrs[] = Libdl.dlsym(handle, @blasfunc(zgetrs_)) + _openblas_cgetrs[] = Libdl.dlsym(handle, @blasfunc(cgetrs_)) + _openblas_dgetrs[] = Libdl.dlsym(handle, @blasfunc(dgetrs_)) + _openblas_sgetrs[] = Libdl.dlsym(handle, @blasfunc(sgetrs_)) + end + return nothing +end + +macro _openblas_function(symbol, pointer) + if VERSION >= v"1.13.0-DEV.0" + return :($(esc(pointer))[]) + end + return :(($(esc(symbol)), libopenblas)) end -function openblas_getrf!( - A::AbstractMatrix{<:ComplexF64}; - ipiv = similar(A, BlasInt, min(size(A, 1), size(A, 2))), - info = Ref{BlasInt}(), - check = false +@inline function openblas_getrf!( + A::AbstractMatrix{<:ComplexF64}, ipiv::AbstractVector{BlasInt}, + info::Ref{BlasInt}, check::Bool ) __openblas_isavailable() || error("Error, OpenBLAS binary is missing but solve is being called. Report this issue") @@ -57,11 +84,10 @@ function openblas_getrf!( chkstride1(A) m, n = size(A) lda = max(1, stride(A, 2)) - if isempty(ipiv) - ipiv = similar(A, BlasInt, min(size(A, 1), size(A, 2))) - end + length(ipiv) == min(m, n) || + throw(DimensionMismatch("ipiv has length $(length(ipiv)), but needs $(min(m, n))")) ccall( - (@blasfunc(zgetrf_), libopenblas), Cvoid, + @_openblas_function(@blasfunc(zgetrf_), _openblas_zgetrf), Cvoid, ( Ref{BlasInt}, Ref{BlasInt}, Ptr{ComplexF64}, Ref{BlasInt}, Ptr{BlasInt}, Ptr{BlasInt}, @@ -69,14 +95,12 @@ function openblas_getrf!( m, n, A, lda, ipiv, info ) chkargsok(info[]) - return A, ipiv, info[], info #Error code is stored in LU factorization type + return info[] end -function openblas_getrf!( - A::AbstractMatrix{<:ComplexF32}; - ipiv = similar(A, BlasInt, min(size(A, 1), size(A, 2))), - info = Ref{BlasInt}(), - check = false +@inline function openblas_getrf!( + A::AbstractMatrix{<:ComplexF32}, ipiv::AbstractVector{BlasInt}, + info::Ref{BlasInt}, check::Bool ) __openblas_isavailable() || error("Error, OpenBLAS binary is missing but solve is being called. Report this issue") @@ -85,11 +109,10 @@ function openblas_getrf!( chkstride1(A) m, n = size(A) lda = max(1, stride(A, 2)) - if isempty(ipiv) - ipiv = similar(A, BlasInt, min(size(A, 1), size(A, 2))) - end + length(ipiv) == min(m, n) || + throw(DimensionMismatch("ipiv has length $(length(ipiv)), but needs $(min(m, n))")) ccall( - (@blasfunc(cgetrf_), libopenblas), Cvoid, + @_openblas_function(@blasfunc(cgetrf_), _openblas_cgetrf), Cvoid, ( Ref{BlasInt}, Ref{BlasInt}, Ptr{ComplexF32}, Ref{BlasInt}, Ptr{BlasInt}, Ptr{BlasInt}, @@ -97,14 +120,12 @@ function openblas_getrf!( m, n, A, lda, ipiv, info ) chkargsok(info[]) - return A, ipiv, info[], info #Error code is stored in LU factorization type + return info[] end -function openblas_getrf!( - A::AbstractMatrix{<:Float64}; - ipiv = similar(A, BlasInt, min(size(A, 1), size(A, 2))), - info = Ref{BlasInt}(), - check = false +@inline function openblas_getrf!( + A::AbstractMatrix{<:Float64}, ipiv::AbstractVector{BlasInt}, + info::Ref{BlasInt}, check::Bool ) __openblas_isavailable() || error("Error, OpenBLAS binary is missing but solve is being called. Report this issue") @@ -113,11 +134,10 @@ function openblas_getrf!( chkstride1(A) m, n = size(A) lda = max(1, stride(A, 2)) - if isempty(ipiv) - ipiv = similar(A, BlasInt, min(size(A, 1), size(A, 2))) - end + length(ipiv) == min(m, n) || + throw(DimensionMismatch("ipiv has length $(length(ipiv)), but needs $(min(m, n))")) ccall( - (@blasfunc(dgetrf_), libopenblas), Cvoid, + @_openblas_function(@blasfunc(dgetrf_), _openblas_dgetrf), Cvoid, ( Ref{BlasInt}, Ref{BlasInt}, Ptr{Float64}, Ref{BlasInt}, Ptr{BlasInt}, Ptr{BlasInt}, @@ -125,14 +145,12 @@ function openblas_getrf!( m, n, A, lda, ipiv, info ) chkargsok(info[]) - return A, ipiv, info[], info #Error code is stored in LU factorization type + return info[] end -function openblas_getrf!( - A::AbstractMatrix{<:Float32}; - ipiv = similar(A, BlasInt, min(size(A, 1), size(A, 2))), - info = Ref{BlasInt}(), - check = false +@inline function openblas_getrf!( + A::AbstractMatrix{<:Float32}, ipiv::AbstractVector{BlasInt}, + info::Ref{BlasInt}, check::Bool ) __openblas_isavailable() || error("Error, OpenBLAS binary is missing but solve is being called. Report this issue") @@ -141,11 +159,10 @@ function openblas_getrf!( chkstride1(A) m, n = size(A) lda = max(1, stride(A, 2)) - if isempty(ipiv) - ipiv = similar(A, BlasInt, min(size(A, 1), size(A, 2))) - end + length(ipiv) == min(m, n) || + throw(DimensionMismatch("ipiv has length $(length(ipiv)), but needs $(min(m, n))")) ccall( - (@blasfunc(sgetrf_), libopenblas), Cvoid, + @_openblas_function(@blasfunc(sgetrf_), _openblas_sgetrf), Cvoid, ( Ref{BlasInt}, Ref{BlasInt}, Ptr{Float32}, Ref{BlasInt}, Ptr{BlasInt}, Ptr{BlasInt}, @@ -153,15 +170,15 @@ function openblas_getrf!( m, n, A, lda, ipiv, info ) chkargsok(info[]) - return A, ipiv, info[], info #Error code is stored in LU factorization type + return info[] end function openblas_getrs!( trans::AbstractChar, A::AbstractMatrix{<:ComplexF64}, ipiv::AbstractVector{BlasInt}, - B::AbstractVecOrMat{<:ComplexF64}; - info = Ref{BlasInt}() + B::AbstractVecOrMat{<:ComplexF64}, + info::Ref{BlasInt} ) __openblas_isavailable() || error("Error, OpenBLAS binary is missing but solve is being called. Report this issue") @@ -177,7 +194,7 @@ function openblas_getrs!( end nrhs = size(B, 2) ccall( - (@blasfunc(zgetrs_), libopenblas), Cvoid, + @_openblas_function(@blasfunc(zgetrs_), _openblas_zgetrs), Cvoid, ( Ref{UInt8}, Ref{BlasInt}, Ref{BlasInt}, Ptr{ComplexF64}, Ref{BlasInt}, Ptr{BlasInt}, Ptr{ComplexF64}, Ref{BlasInt}, Ptr{BlasInt}, Clong, @@ -193,8 +210,8 @@ function openblas_getrs!( trans::AbstractChar, A::AbstractMatrix{<:ComplexF32}, ipiv::AbstractVector{BlasInt}, - B::AbstractVecOrMat{<:ComplexF32}; - info = Ref{BlasInt}() + B::AbstractVecOrMat{<:ComplexF32}, + info::Ref{BlasInt} ) __openblas_isavailable() || error("Error, OpenBLAS binary is missing but solve is being called. Report this issue") @@ -210,7 +227,7 @@ function openblas_getrs!( end nrhs = size(B, 2) ccall( - (@blasfunc(cgetrs_), libopenblas), Cvoid, + @_openblas_function(@blasfunc(cgetrs_), _openblas_cgetrs), Cvoid, ( Ref{UInt8}, Ref{BlasInt}, Ref{BlasInt}, Ptr{ComplexF32}, Ref{BlasInt}, Ptr{BlasInt}, Ptr{ComplexF32}, Ref{BlasInt}, Ptr{BlasInt}, Clong, @@ -226,8 +243,8 @@ function openblas_getrs!( trans::AbstractChar, A::AbstractMatrix{<:Float64}, ipiv::AbstractVector{BlasInt}, - B::AbstractVecOrMat{<:Float64}; - info = Ref{BlasInt}() + B::AbstractVecOrMat{<:Float64}, + info::Ref{BlasInt} ) __openblas_isavailable() || error("Error, OpenBLAS binary is missing but solve is being called. Report this issue") @@ -243,7 +260,7 @@ function openblas_getrs!( end nrhs = size(B, 2) ccall( - (@blasfunc(dgetrs_), libopenblas), Cvoid, + @_openblas_function(@blasfunc(dgetrs_), _openblas_dgetrs), Cvoid, ( Ref{UInt8}, Ref{BlasInt}, Ref{BlasInt}, Ptr{Float64}, Ref{BlasInt}, Ptr{BlasInt}, Ptr{Float64}, Ref{BlasInt}, Ptr{BlasInt}, Clong, @@ -259,8 +276,8 @@ function openblas_getrs!( trans::AbstractChar, A::AbstractMatrix{<:Float32}, ipiv::AbstractVector{BlasInt}, - B::AbstractVecOrMat{<:Float32}; - info = Ref{BlasInt}() + B::AbstractVecOrMat{<:Float32}, + info::Ref{BlasInt} ) __openblas_isavailable() || error("Error, OpenBLAS binary is missing but solve is being called. Report this issue") @@ -276,7 +293,7 @@ function openblas_getrs!( end nrhs = size(B, 2) ccall( - (@blasfunc(sgetrs_), libopenblas), Cvoid, + @_openblas_function(@blasfunc(sgetrs_), _openblas_sgetrs), Cvoid, ( Ref{UInt8}, Ref{BlasInt}, Ref{BlasInt}, Ptr{Float32}, Ref{BlasInt}, Ptr{BlasInt}, Ptr{Float32}, Ref{BlasInt}, Ptr{BlasInt}, Clong, @@ -293,9 +310,28 @@ _get_residualsafety(alg::OpenBLASLUFactorization) = alg.residualsafety default_alias_A(::OpenBLASLUFactorization, ::Any, ::Any) = false default_alias_b(::OpenBLASLUFactorization, ::Any, ::Any) = false -const PREALLOCATED_OPENBLAS_LU = begin - A = rand(0, 0) - luinst = ArrayInterface.lu_instance(A), Ref{BlasInt}() +mutable struct OpenBLASLUCache{F, P, I} + factors::F + ipiv::P + info::I +end + +_custom_cache_factorization(::OpenBLASLUFactorization, cacheval::OpenBLASLUCache) = + LU(cacheval.factors, cacheval.ipiv, Int(cacheval.info[])) + +@inline function _direct_lu_factorize!( + cacheval::OpenBLASLUCache, A_work, ::OpenBLASLUFactorization + ) + cacheval.factors = A_work + return openblas_getrf!(A_work, cacheval.ipiv, cacheval.info, false) +end + +@inline function _direct_lu_solve!( + cacheval::OpenBLASLUCache, u, b, ::OpenBLASLUFactorization + ) + copyto!(u, b) + openblas_getrs!('N', cacheval.factors, cacheval.ipiv, u, cacheval.info) + return u end function LinearSolve.init_cacheval( @@ -303,7 +339,8 @@ function LinearSolve.init_cacheval( maxiters::Int, abstol, reltol, verbose::Union{LinearVerbosity, Bool}, assumptions::OperatorAssumptions ) - return PREALLOCATED_OPENBLAS_LU + A0 = Matrix{Float64}(undef, 0, 0) + return OpenBLASLUCache(A0, Vector{BlasInt}(undef, 0), Ref{BlasInt}()) end function LinearSolve.init_cacheval( @@ -312,9 +349,9 @@ function LinearSolve.init_cacheval( maxiters::Int, abstol, reltol, verbose::Union{LinearVerbosity, Bool}, assumptions::OperatorAssumptions ) - # Ask `lu_instance` about `A` itself so wrapper-specific dispatch can choose - # the factorization container that `solve!` will store. - return ArrayInterface.lu_instance(A), Ref{BlasInt}() + return OpenBLASLUCache( + A, similar(A, BlasInt, min(size(A, 1), size(A, 2))), Ref{BlasInt}() + ) end function SciMLBase.solve!( @@ -323,60 +360,64 @@ function SciMLBase.solve!( ) __openblas_isavailable() || error("Error, OpenBLAS binary is missing but solve is being called. Report this issue") - A = cache.A - A = convert(AbstractMatrix, A) + A_work = convert(AbstractMatrix, cache.A) check_safety = alg.residualsafety && cache.isfresh needs_backup = check_safety || (cache.alg isa DefaultLinearSolver && cache.alg.safetyfallback && cache.isfresh) - A_original = needs_backup ? _copy_A_for_safety(cache) : A + A_original = needs_backup ? _copy_A_for_safety(cache) : A_work verbose = cache.verbose if cache.isfresh cacheval = @get_cacheval(cache, :OpenBLASLUFactorization) - res = openblas_getrf!(A; ipiv = cacheval[1].ipiv, info = cacheval[2]) - fact = LU(res[1:3]...), res[4] - cache.cacheval = fact - - info_value = res[3] + if length(cacheval.ipiv) != min(size(A_work, 1), size(A_work, 2)) + cacheval.ipiv = similar( + A_work, BlasInt, min(size(A_work, 1), size(A_work, 2)) + ) + end + info_value = _direct_lu_factorize!(cacheval, A_work, alg) if info_value != 0 if verbose.blas_info != SciMLLogging.Silent() || verbose.blas_errors != SciMLLogging.Silent() || verbose.blas_invalid_args != SciMLLogging.Silent() - op_info = get_blas_operation_info( - :dgetrf, A, cache.b, + failure_op_info = get_blas_operation_info( + :dgetrf, A_work, cache.b, condition = verbose.condition_number != SciMLLogging.Silent() ) - @SciMLMessage(cache.verbose, :condition_number) do - if isinf(op_info.condition_number) - return "Matrix condition number calculation failed." - else - return "Matrix condition number: $(round(op_info.condition_number, sigdigits = 4)) for $(size(A, 1))×$(size(A, 2)) matrix in dgetrf" + let op_info = failure_op_info + @SciMLMessage(cache.verbose, :condition_number) do + if isinf(op_info.condition_number) + return "Matrix condition number calculation failed." + else + return "Matrix condition number: $(round(op_info.condition_number, sigdigits = 4)) for $(size(A_work, 1))×$(size(A_work, 2)) matrix in dgetrf" + end end end verb_option, message = blas_info_msg( - :dgetrf, info_value; extra_context = op_info + :dgetrf, info_value; extra_context = failure_op_info ) @SciMLMessage(message, verbose, verb_option) end else @SciMLMessage(cache.verbose, :blas_success) do - op_info = get_blas_operation_info( - :dgetrf, A, cache.b, + success_op_info = get_blas_operation_info( + :dgetrf, A_work, cache.b, condition = verbose.condition_number != SciMLLogging.Silent() ) - @SciMLMessage(cache.verbose, :condition_number) do - if isinf(op_info.condition_number) - return "Matrix condition number calculation failed." - else - return "Matrix condition number: $(round(op_info.condition_number, sigdigits = 4)) for $(size(A, 1))×$(size(A, 2)) matrix in dgetrf" + let op_info = success_op_info + @SciMLMessage(cache.verbose, :condition_number) do + if isinf(op_info.condition_number) + return "Matrix condition number calculation failed." + else + return "Matrix condition number: $(round(op_info.condition_number, sigdigits = 4)) for $(size(A_work, 1))×$(size(A_work, 2)) matrix in dgetrf" + end end end - return "BLAS LU factorization (dgetrf) completed successfully for $(op_info.matrix_size) matrix" + return "BLAS LU factorization (dgetrf) completed successfully for $(success_op_info.matrix_size) matrix" end end - if !LinearAlgebra.issuccess(fact[1]) + if info_value != 0 @SciMLMessage("Solver failed", cache.verbose, :solver_failure) return SciMLBase.build_linear_solution( alg, cache.u, nothing, nothing; retcode = ReturnCode.Failure @@ -385,20 +426,21 @@ function SciMLBase.solve!( cache.isfresh = false end - A, info = @get_cacheval(cache, :OpenBLASLUFactorization) + cacheval = @get_cacheval(cache, :OpenBLASLUFactorization) + factors = cacheval.factors + info = cacheval.info require_one_based_indexing(cache.u, cache.b) - m, n = size(A, 1), size(A, 2) + m, n = size(factors, 1), size(factors, 2) if m > n Bc = copy(cache.b) - openblas_getrs!('N', A.factors, A.ipiv, Bc; info) + openblas_getrs!('N', factors, cacheval.ipiv, Bc, info) if cache.b isa AbstractMatrix copyto!(cache.u, @view(Bc[1:n, :])) else copyto!(cache.u, 1, Bc, 1, n) end else - copyto!(cache.u, cache.b) - openblas_getrs!('N', A.factors, A.ipiv, cache.u; info) + _direct_lu_solve!(cacheval, cache.u, cache.b, alg) end if check_safety @@ -431,7 +473,8 @@ function LinearSolve.init_cacheval( A_32 = similar(A, T32) b_32 = similar(b, T32) u_32 = similar(u, T32) - luinst = ArrayInterface.lu_instance(rand(T32, 0, 0)) + ipiv = similar(A_32, BlasInt, min(size(A_32, 1), size(A_32, 2))) + luinst = LU(A_32, ipiv, zero(BlasInt)) # Return tuple with pre-allocated arrays return (luinst, Ref{BlasInt}(), A_32, b_32, u_32) end @@ -452,11 +495,9 @@ function SciMLBase.solve!( # Compute 32-bit type on demand and copy A T32 = eltype(A) <: Complex ? ComplexF32 : Float32 A_32 .= T32.(A) - res = openblas_getrf!(A_32; ipiv = luinst.ipiv, info = info) - fact = (LU(res[1:3]...), res[4], A_32, b_32, u_32) - cache.cacheval = fact + info_value = openblas_getrf!(A_32, luinst.ipiv, info, false) - if !LinearAlgebra.issuccess(fact[1]) + if info_value != 0 @SciMLMessage("Solver failed", cache.verbose, :solver_failure) return SciMLBase.build_linear_solution( alg, cache.u, nothing, nothing; retcode = ReturnCode.Failure @@ -477,7 +518,7 @@ function SciMLBase.solve!( b_32 .= T32.(cache.b) if m > n - openblas_getrs!('N', A_lu.factors, A_lu.ipiv, b_32; info) + openblas_getrs!('N', A_lu.factors, A_lu.ipiv, b_32, info) # Convert back to original precision if cache.b isa AbstractMatrix cache.u .= Torig.(@view(b_32[1:n, :])) @@ -486,7 +527,7 @@ function SciMLBase.solve!( end else copyto!(u_32, b_32) - openblas_getrs!('N', A_lu.factors, A_lu.ipiv, u_32; info) + openblas_getrs!('N', A_lu.factors, A_lu.ipiv, u_32, info) # Convert back to original precision cache.u .= Torig.(u_32) end diff --git a/test/AD/enzyme.jl b/test/AD/enzyme.jl index 141d565af..a3802f502 100644 --- a/test/AD/enzyme.jl +++ b/test/AD/enzyme.jl @@ -27,6 +27,31 @@ db12 = ForwardDiff.gradient(x -> f(eltype(x).(A), x), copy(b1)) @test dA ≈ dA2 @test db1 ≈ db12 +@testset "OpenBLAS solve! reverse rule" begin + A = [3.0 1.0; 1.0 2.0] + b = [1.0, 2.0] + dA = zeros(size(A)) + db = zeros(size(b)) + + function openblas_solve!(A, b) + cache = init(LinearProblem(A, b), OpenBLASLUFactorization()) + return sum(solve!(cache).u) + end + + Enzyme.autodiff( + Reverse, openblas_solve!, Duplicated(copy(A), dA), Duplicated(copy(b), db) + ) + expected_dA = FiniteDiff.finite_difference_gradient( + A -> openblas_solve!(A, b), A + ) + expected_db = FiniteDiff.finite_difference_gradient( + b -> openblas_solve!(A, b), b + ) + + @test dA ≈ expected_dA + @test db ≈ expected_db +end + A = rand(n, n); dA = zeros(n, n); b1 = rand(n); diff --git a/test/AD/mooncake.jl b/test/AD/mooncake.jl index 0113a1c27..9f779bc06 100644 --- a/test/AD/mooncake.jl +++ b/test/AD/mooncake.jl @@ -3,6 +3,26 @@ using LinearSolve, LinearAlgebra, Test using FiniteDiff, RecursiveFactorization using Mooncake +struct OpaqueDenseFactorization <: LinearSolve.AbstractDenseFactorization end +struct OpaqueDenseFactorizationCache end + +function LinearSolve.init_cacheval( + ::OpaqueDenseFactorization, A, b, u, Pl, Pr, maxiters::Int, + abstol, reltol, verbose::Union{LinearSolve.LinearVerbosity, Bool}, + assumptions::LinearSolve.OperatorAssumptions + ) + return OpaqueDenseFactorizationCache() +end + +function LinearSolve.solve!( + cache::LinearSolve.LinearCache, alg::OpaqueDenseFactorization; kwargs... + ) + cache.u .= cache.A \ cache.b + return LinearSolve.SciMLBase.build_linear_solution( + alg, cache.u, nothing, nothing; retcode = ReturnCode.Success + ) +end + # first test n = 4 A = rand(n, n); @@ -28,6 +48,45 @@ db12 = ForwardDiff.gradient(x -> f(eltype(x).(A), x), copy(b1)) @test gradient[2] ≈ dA2 @test gradient[3] ≈ db12 +@testset "OpenBLAS solve! reverse rule" begin + A = [3.0 1.0; 1.0 2.0] + b = [1.0, 2.0] + + function openblas_solve!(A, b) + cache = init(LinearProblem(A, b), OpenBLASLUFactorization()) + return sum(solve!(cache).u) + end + + rule = Mooncake.build_rrule(openblas_solve!, copy(A), copy(b)) + value, gradient = Mooncake.value_and_pullback!!( + rule, 1.0, openblas_solve!, copy(A), copy(b) + ) + dA = FiniteDiff.finite_difference_gradient(A -> openblas_solve!(A, b), A) + db = FiniteDiff.finite_difference_gradient(b -> openblas_solve!(A, b), b) + + @test value ≈ openblas_solve!(A, b) + @test gradient[2] ≈ dA + @test gradient[3] ≈ db +end + +@testset "Opaque factorization cache solve! reverse rule" begin + A = rand(4, 4) + 4I + b = rand(4) + + function opaque_factorization_solve!(b) + cache = init(LinearProblem(A, b), OpaqueDenseFactorization()) + return sum(solve!(cache).u) + end + + rule = Mooncake.build_rrule(opaque_factorization_solve!, copy(b)) + value, gradient = Mooncake.value_and_pullback!!( + rule, 1.0, opaque_factorization_solve!, copy(b) + ) + + @test value ≈ opaque_factorization_solve!(b) + @test gradient[2] ≈ transpose(A) \ ones(4) +end + # Second test A = rand(n, n); b1 = rand(n); diff --git a/test/Core/adjoint.jl b/test/Core/adjoint.jl index 3183e43ef..4c4d5ea7b 100644 --- a/test/Core/adjoint.jl +++ b/test/Core/adjoint.jl @@ -1,13 +1,143 @@ using Zygote, ForwardDiff using LinearSolve, LinearAlgebra, Test using FiniteDiff, RecursiveFactorization -using Random +using InteractiveUtils, Random, SparseArrays +import CliqueTrees + +struct UnregisteredFactorization <: LinearSolve.AbstractFactorization end + +function factorization_leaf_types(T) + children = subtypes(T) + isempty(children) && return Any[T] + return reduce(vcat, factorization_leaf_types.(children); init = Any[]) +end + +if Sys.islinux() + import LAPACK_jll, blis_jll +end Random.seed!(1234) n = 4 A = rand(n, n); b1 = rand(n); +@testset "Adjoint factorization cache dispatch follows the solver" begin + factorization = lu(copy(A)) + @test LinearSolve._cache_factorization(LUFactorization(), factorization) === + factorization + @test LinearSolve._cache_factorization( + GenericLUFactorization(), (factorization, factorization.ipiv) + ) === factorization + @test LinearSolve._can_reuse_cache_factorization( + LUFactorization(), factorization + ) + + krylov = KrylovJL_GMRES() + @test isnothing(LinearSolve._cache_factorization(krylov, factorization)) + @test !LinearSolve._can_reuse_cache_factorization(krylov, factorization) + + default = LinearSolve.DefaultLinearSolver( + LinearSolve.DefaultAlgorithmChoice.LUFactorization + ) + @test isnothing(LinearSolve._cache_factorization(default, factorization)) + @test !LinearSolve._can_reuse_cache_factorization(default, factorization) + + unregistered = UnregisteredFactorization() + @test isnothing(LinearSolve._cache_factorization(unregistered, factorization)) + @test !LinearSolve._can_reuse_cache_factorization(unregistered, factorization) +end + +@testset "Complex Krylov adjoint solve" begin + A_complex = ComplexF64[3 + 1im 1 - 2im; 2 + 0.5im 4 - 1im] + rhs_complex = ComplexF64[0.7 + 0.2im, -0.3 + 0.4im] + adjoint_solution = LinearSolve._adjoint_krylov_solve( + KrylovJL_GMRES(), A_complex, rhs_complex; + abstol = 1.0e-12, reltol = 1.0e-12, verbose = false + ) + @test adjoint(A_complex) * adjoint_solution ≈ rhs_complex +end + +@testset "Every factorization algorithm declares its adjoint reuse policy" begin + for T in factorization_leaf_types(LinearSolve.AbstractFactorization) + parentmodule(T) === LinearSolve || continue + reuse = LinearSolve._adjoint_factorization_reuse(T) + @test !(reuse isa LinearSolve._UnspecifiedAdjointFactorizationReuse) + end +end + +@testset "Solver-specific cached adjoint solves" begin + A_local = [4.0 1.0; 2.0 3.0] + b_local = [1.0, 2.0] + adjoint_rhs = [0.7, -0.3] + + for alg in ( + NormalCholeskyFactorization(), + NormalBunchKaufmanFactorization(), + SimpleLUFactorization(), + ) + cache = init(LinearProblem(copy(A_local), copy(b_local)), alg) + @test LinearSolve._can_reuse_cache_factorization(alg, cache.cacheval) + solve!(cache) + adjoint_solution = LinearSolve._adjoint_factorization_solve( + alg, cache.cacheval, cache.A, adjoint_rhs + ) + @test adjoint(A_local) * adjoint_solution ≈ adjoint_rhs + end + + tall_A = [4.0 1.0; 2.0 3.0; 1.0 -1.0] + tall_b = [1.0, 2.0, -0.5] + tall_adjoint_rhs = [0.7, -0.3] + for alg in (NormalCholeskyFactorization(), NormalBunchKaufmanFactorization()) + cache = init(LinearProblem(copy(tall_A), copy(tall_b)), alg) + solve!(cache) + adjoint_solution = LinearSolve._adjoint_factorization_solve( + alg, cache.cacheval, cache.A, tall_adjoint_rhs + ) + @test adjoint_solution ≈ tall_A * ((adjoint(tall_A) * tall_A) \ tall_adjoint_rhs) + @test adjoint(tall_A) * adjoint_solution ≈ tall_adjoint_rhs + end + + sparse_A = sparse(A_local) + sparse_alg = SparseColumnPivotedQRFactorization() + sparse_cache = init(LinearProblem(copy(sparse_A), copy(b_local)), sparse_alg) + @test LinearSolve._can_reuse_cache_factorization( + sparse_alg, sparse_cache.cacheval + ) + solve!(sparse_cache) + sparse_adjoint_solution = LinearSolve._adjoint_factorization_solve( + sparse_alg, sparse_cache.cacheval, sparse_cache.A, adjoint_rhs + ) + @test adjoint(sparse_A) * sparse_adjoint_solution ≈ adjoint_rhs + + clique_A = sparse([4.0 1.0; 1.0 3.0]) + clique_alg = CliqueTreesFactorization() + clique_cache = init(LinearProblem(copy(clique_A), copy(b_local)), clique_alg) + @test LinearSolve._can_reuse_cache_factorization( + clique_alg, clique_cache.cacheval + ) + solve!(clique_cache) + clique_adjoint_solution = LinearSolve._adjoint_factorization_solve( + clique_alg, clique_cache.cacheval, clique_cache.A, adjoint_rhs + ) + @test adjoint(clique_A) * clique_adjoint_solution ≈ adjoint_rhs + + for alg in (NormalCholeskyFactorization(), SimpleLUFactorization(), sparse_alg) + A_alg = alg isa SparseColumnPivotedQRFactorization ? sparse_A : A_local + db, = Zygote.gradient(b -> sum(solve(LinearProblem(A_alg, b), alg).u), b_local) + @test db ≈ adjoint(A_local) \ ones(2) + end +end + +@testset "Uncached factorization adjoint fallback" begin + diagonal = rand(n) .+ 1 + A_diagonal = Diagonal(diagonal) + f_diagonal(b) = sum( + solve(LinearProblem(A_diagonal, b), DiagonalFactorization()).u + ) + db, = Zygote.gradient(f_diagonal, b1) + @test db ≈ inv.(diagonal) +end + function f(A, b1; alg = LUFactorization()) prob = LinearProblem(A, b1) @@ -103,11 +233,16 @@ db22 = ForwardDiff.gradient(x -> f4(eltype(x).(A), eltype(x).(b1), x), copy(b1)) A = rand(n, n); b1 = rand(n); -for alg in ( - LUFactorization(), - RFLUFactorization(), - KrylovJL_GMRES(), - ) +adjoint_algs = Any[ + LUFactorization(), + RFLUFactorization(), + KrylovJL_GMRES(), +] +LinearSolve.useopenblas && push!(adjoint_algs, OpenBLASLUFactorization()) +if Base.get_extension(LinearSolve, :LinearSolveBLISExt) !== nothing + push!(adjoint_algs, LinearSolve.BLISLUFactorization()) +end +for alg in adjoint_algs @show alg function fb(b) prob = LinearProblem(A, b) diff --git a/test/Core/basictests.jl b/test/Core/basictests.jl index 3176893a2..f97f0be2e 100644 --- a/test/Core/basictests.jl +++ b/test/Core/basictests.jl @@ -439,7 +439,7 @@ end # Test BLIS if extension is available if Base.get_extension(LinearSolve, :LinearSolveBLISExt) !== nothing - push!(test_algs, BLISLUFactorization()) + push!(test_algs, LinearSolve.BLISLUFactorization()) end @testset "Concrete Factorizations" begin diff --git a/test/Core/butterfly.jl b/test/Core/butterfly.jl index b9083171c..5189f5617 100644 --- a/test/Core/butterfly.jl +++ b/test/Core/butterfly.jl @@ -12,6 +12,20 @@ using RecursiveFactorization end end +@testset "Cached adjoint solve" begin + n = 16 + A = rand(n, n) + b = rand(n) + cache = init(LinearProblem(A, b), ButterflyFactorization()) + @test LinearSolve._can_reuse_cache_factorization(cache.alg, cache.cacheval) + solve!(cache) + adjoint_rhs = rand(n) + adjoint_solution = LinearSolve._adjoint_factorization_solve( + cache.alg, cache.cacheval, cache.A, adjoint_rhs + ) + @test adjoint(A) * adjoint_solution ≈ adjoint_rhs +end + function wilkinson(N) A = zeros(N, N) A[1:(N + 1):(N * N)] .= 1 diff --git a/test/Core/componentarrays.jl b/test/Core/componentarrays.jl index bcded4592..d33dff9c2 100644 --- a/test/Core/componentarrays.jl +++ b/test/Core/componentarrays.jl @@ -27,7 +27,8 @@ using LinearSolve: OperatorAssumptions true, assump, ) - @test slot[1].factors isa ComponentMatrix + workspace = alg isa MKLLUFactorization ? slot[1] : slot + @test workspace.factors isa ComponentMatrix end slot = LinearSolve.init_cacheval( @@ -43,8 +44,8 @@ using LinearSolve: OperatorAssumptions true, assump, ) - @test slot[1].factors isa ComponentMatrix - @test slot[1].ipiv isa Vector{Cint} + @test slot.factors isa ComponentMatrix + @test slot.ipiv isa Vector{Cint} end @testset "dense ComponentMatrix default solve" begin diff --git a/test/Core/direct_blas_refactorization.jl b/test/Core/direct_blas_refactorization.jl new file mode 100644 index 000000000..88f0a431f --- /dev/null +++ b/test/Core/direct_blas_refactorization.jl @@ -0,0 +1,91 @@ +using LinearSolve, LinearAlgebra, StableRNGs, Test + +if Sys.islinux() + import LAPACK_jll, blis_jll +end + +const direct_blas_rng = StableRNG(42) + +function direct_blas_refactor_solve!(cache, Awork, A) + copyto!(Awork, A) + cache.A = Awork + return solve!(cache) +end + +function test_direct_blas_refactorization(alg, ::Type{T}) where {T} + n = 51 + A1 = rand(direct_blas_rng, T, n, n) + n * I + A2 = rand(direct_blas_rng, T, n, n) + n * I + Asing = copy(A1) + Asing[:, 1] .= zero(T) + b = rand(direct_blas_rng, T, n) + cache = init(LinearProblem(copy(A1), copy(b)), alg) + Awork = cache.A + + @test solve!(cache).u ≈ A1 \ b + for Ak in (A2, A1, A2) + sol = direct_blas_refactor_solve!(cache, Awork, Ak) + @test sol.retcode == ReturnCode.Success + @test sol.u ≈ Ak \ b + end + + direct_blas_refactor_solve!(cache, Awork, A1) + ipiv_before = cache.cacheval.ipiv + direct_blas_refactor_solve!(cache, Awork, A2) + @test cache.cacheval.ipiv === ipiv_before + @test cache.cacheval.factors === Awork + @test cache.u ≈ A2 \ b + @test LinearSolve._cache_factorization(alg, cache.cacheval) \ b ≈ A2 \ b + @test LinearSolve._can_reuse_cache_factorization(alg, cache.cacheval) + + @test direct_blas_refactor_solve!(cache, Awork, Asing).retcode == + ReturnCode.Failure + sol = direct_blas_refactor_solve!(cache, Awork, A1) + @test sol.retcode == ReturnCode.Success + return @test sol.u ≈ A1 \ b +end + +function test_direct_blas_resize(alg) + n1, n2 = 5, 9 + A1 = rand(direct_blas_rng, n1, n1) + n1 * I + b1 = rand(direct_blas_rng, n1) + cache = init(LinearProblem(copy(A1), copy(b1)), alg) + @test solve!(cache).u ≈ A1 \ b1 + ipiv_before = cache.cacheval.ipiv + + resize!(cache, n2) + A2 = rand(direct_blas_rng, n2, n2) + n2 * I + b2 = rand(direct_blas_rng, n2) + cache.A = copy(A2) + cache.b = copy(b2) + cache.u = zeros(n2) + @test solve!(cache).u ≈ A2 \ b2 + @test length(cache.cacheval.ipiv) == n2 + @test cache.cacheval.ipiv !== ipiv_before + + Awork = cache.A + direct_blas_refactor_solve!(cache, Awork, A2) + return @test direct_blas_refactor_solve!(cache, Awork, A2).u ≈ A2 \ b2 +end + +if LinearSolve.useopenblas + @testset "OpenBLAS reuses its direct LU workspace" begin + for T in (Float32, Float64, ComplexF32, ComplexF64) + @testset "$T" test_direct_blas_refactorization( + OpenBLASLUFactorization(), T + ) + end + test_direct_blas_resize(OpenBLASLUFactorization()) + end +end + +if Base.get_extension(LinearSolve, :LinearSolveBLISExt) !== nothing + @testset "BLIS reuses its direct LU workspace" begin + for T in (Float32, Float64, ComplexF32, ComplexF64) + @testset "$T" test_direct_blas_refactorization( + LinearSolve.BLISLUFactorization(), T + ) + end + test_direct_blas_resize(LinearSolve.BLISLUFactorization()) + end +end diff --git a/test/Core/fixedsizearrays.jl b/test/Core/fixedsizearrays.jl index 14b1f3ddc..5e3c301da 100644 --- a/test/Core/fixedsizearrays.jl +++ b/test/Core/fixedsizearrays.jl @@ -83,10 +83,10 @@ end end # The BLAS-direct LU caches (`MKL`, `OpenBLAS`, `AppleAccelerate`, `BLIS`) store -# a factorization built from `A` into a type-parameterized `cacheval` slot, so -# `init_cacheval` must produce a slot whose container matches `A`. This holds -# regardless of whether the corresponding BLAS binary is present, so the type -# check runs everywhere even when the solver itself can't. +# factorization buffers built from `A` in a type-parameterized `cacheval` slot, +# so `init_cacheval` must produce buffers whose container matches `A`. This +# holds regardless of whether the corresponding BLAS binary is present, so the +# type check runs everywhere even when the solver itself can't. @testset "BLAS LU init_cacheval slot tracks the FixedSizeArray container" begin n = 20 A = FixedSizeArray(rand(n, n) + n * I) @@ -99,7 +99,8 @@ end slot = LinearSolve.init_cacheval( alg, A, v, v, nothing, nothing, 0, 0.0, 0.0, true, assump ) - @test slot[1].factors isa FixedSizeArray - @test slot[1].ipiv isa FixedSizeArray + workspace = alg isa MKLLUFactorization ? slot[1] : slot + @test workspace.factors isa FixedSizeArray + @test workspace.ipiv isa FixedSizeArray end end diff --git a/test/Core/lu_refactorization.jl b/test/Core/lu_refactorization.jl index 736b3a3d7..d0ec6c286 100644 --- a/test/Core/lu_refactorization.jl +++ b/test/Core/lu_refactorization.jl @@ -109,6 +109,44 @@ end end end +@testset "Apple factorization views use the stdlib pivot type" begin + A = [4.0 1.0; 2.0 3.0] + b = [1.0, 2.0] + fact = lu(A) + cacheval = LinearSolve.AppleAccelerateLUCache( + fact.factors, Cint.(fact.ipiv), Ref{Cint}(Cint(fact.info)) + ) + alg = AppleAccelerateLUFactorization() + view = LinearSolve._cache_factorization(alg, cacheval) + @test eltype(view.ipiv) === LinearAlgebra.BlasInt + @test view \ b ≈ A \ b + @test LinearSolve._can_reuse_cache_factorization(alg, cacheval) +end + +if LinearSolve.appleaccelerate_isavailable() + @testset "Apple Accelerate reuses its refactorization workspace" begin + n = 51 + A1 = rand(rng, n, n) + n * I + A2 = rand(rng, n, n) + n * I + Asing = copy(A1) + Asing[:, 1] .= 0 + b = rand(rng, n) + cache = init(LinearProblem(copy(A1), copy(b)), AppleAccelerateLUFactorization()) + Awork = cache.A + + @test solve!(cache).u ≈ A1 \ b + refactor_solve!(cache, Awork, A1) + refactor_solve!(cache, Awork, A2) + @test cache.u ≈ A2 \ b + @test LinearSolve._cache_factorization(cache.alg, cache.cacheval) \ b ≈ A2 \ b + + @test refactor_solve!(cache, Awork, Asing).retcode == ReturnCode.Failure + sol = refactor_solve!(cache, Awork, A1) + @test sol.retcode == ReturnCode.Success + @test sol.u ≈ A1 \ b + end +end + @testset "default algorithm QR safety fallback survives warm singular refactorization" begin n = 8 Agood = rand(rng, n, n) + n * I diff --git a/test/Core/verbosity.jl b/test/Core/verbosity.jl index 13535b14e..6005091d9 100644 --- a/test/Core/verbosity.jl +++ b/test/Core/verbosity.jl @@ -240,7 +240,9 @@ end blas_info = Silent() ) - @test_logs (:warn, r"BLAS/LAPACK.*Matrix is singular") solve( + @test_logs ( + :warn, r"BLAS/LAPACK.*Matrix is singular", + ) (:warn, r"Solver failed") solve( prob_singular, BLISLUFactorization(); verbose = verbose_errors ) @@ -251,7 +253,9 @@ end blas_success = Silent() ) - @test_logs (:info, r"BLAS/LAPACK.*Matrix is singular") solve( + @test_logs ( + :info, r"BLAS/LAPACK.*Matrix is singular", + ) (:warn, r"Solver failed") solve( prob_singular, BLISLUFactorization(); verbose = verbose_info ) diff --git a/test/LinearSolveHSL/hsl.jl b/test/LinearSolveHSL/hsl.jl index b954d36f8..ee9c52686 100644 --- a/test/LinearSolveHSL/hsl.jl +++ b/test/LinearSolveHSL/hsl.jl @@ -40,6 +40,11 @@ else @test sol1.retcode == ReturnCode.Success @test sol2.retcode == ReturnCode.Success @test A * sol2.u ≈ b2 rtol = 1.0e-9 atol = 1.0e-11 + adjoint_rhs = rand(Float64, n) + adjoint_solution = LinearSolve._adjoint_factorization_solve( + cache.alg, cache.cacheval, cache.A, adjoint_rhs + ) + @test adjoint(A) * adjoint_solution ≈ adjoint_rhs rtol = 1.0e-9 atol = 1.0e-11 end @testset "HSL MA97 wrapper" begin @@ -63,6 +68,11 @@ else @test sol1.retcode == ReturnCode.Success @test sol2.retcode == ReturnCode.Success @test A * sol2.u ≈ b2 rtol = 1.0e-9 atol = 1.0e-11 + adjoint_rhs = rand(Float64, n) + adjoint_solution = LinearSolve._adjoint_factorization_solve( + cache.alg, cache.cacheval, cache.A, adjoint_rhs + ) + @test adjoint(A) * adjoint_solution ≈ adjoint_rhs rtol = 1.0e-9 atol = 1.0e-11 Ac = ComplexF64.(A) bc = rand(ComplexF64, n) diff --git a/test/LinearSolveMUMPS/mumps.jl b/test/LinearSolveMUMPS/mumps.jl index edfbd54c7..f1e4253ea 100644 --- a/test/LinearSolveMUMPS/mumps.jl +++ b/test/LinearSolveMUMPS/mumps.jl @@ -51,6 +51,11 @@ end sol3 = LinearSolve.solve!(cache) @test sol3.retcode == ReturnCode.Success @test residual_ok(A2, sol3.u, b2) + adjoint_rhs = [0.7, -0.3, 0.2] + adjoint_solution = LinearSolve._adjoint_factorization_solve( + alg, cache.cacheval, cache.A, adjoint_rhs + ) + @test adjoint(A2) * adjoint_solution ≈ adjoint_rhs MUMPSExt.cleanup_mumps_cache!(cache) end @@ -71,6 +76,11 @@ end sol = LinearSolve.solve!(cache) @test sol.retcode == ReturnCode.Success @test sol.u ≈ x atol = 1.0e-8 rtol = 1.0e-8 + adjoint_rhs = ComplexF64[0.7 + 0.2im, -0.3 + 0.4im] + adjoint_solution = LinearSolve._adjoint_factorization_solve( + cache.alg, cache.cacheval, cache.A, adjoint_rhs + ) + @test adjoint(transpose(A)) * adjoint_solution ≈ adjoint_rhs MUMPSExt.cleanup_mumps_cache!(cache) end diff --git a/test/LinearSolvePardiso/pardiso.jl b/test/LinearSolvePardiso/pardiso.jl index 8ff7c043b..1ecb20227 100644 --- a/test/LinearSolvePardiso/pardiso.jl +++ b/test/LinearSolvePardiso/pardiso.jl @@ -75,6 +75,21 @@ for alg in algs @test sol11.u ≈ sol31.u @test sol12.u ≈ sol32.u @test sol13.u ≈ sol33.u + adjoint_rhs = rand(n) + adjoint_solution = LinearSolve._adjoint_factorization_solve( + alg, linsolve.cacheval, linsolve.A, adjoint_rhs + ) + @test adjoint(A2) * adjoint_solution ≈ adjoint_rhs + + complex_A = complex.(A2, sprand(n, n, 0.1)) + complex_b = rand(ComplexF64, n) + complex_cache = init(LinearProblem(complex_A, complex_b), alg) + solve!(complex_cache) + complex_adjoint_rhs = rand(ComplexF64, n) + complex_adjoint_solution = LinearSolve._adjoint_factorization_solve( + alg, complex_cache.cacheval, complex_cache.A, complex_adjoint_rhs + ) + @test adjoint(complex_A) * complex_adjoint_solution ≈ complex_adjoint_rhs end # Test for problem from #497 diff --git a/test/qa/Project.toml b/test/qa/Project.toml index 13a926122..253c5645c 100644 --- a/test/qa/Project.toml +++ b/test/qa/Project.toml @@ -1,19 +1,27 @@ [deps] +AllocCheck = "9b6a8646-10ed-4001-bbdc-1d2f46dfbb1a" Aqua = "4c88cf16-eb10-579e-8560-4a9242c79595" JET = "c3a54625-cd67-489e-a8e7-0a5a0ff4e31b" +LAPACK_jll = "51474c39-65e3-53ba-86ba-03b1b862ec14" +LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e" LinearSolve = "7ed4a6bd-45f5-4d41-b270-4a48e9bafcae" SafeTestsets = "1bc83da4-3b8d-516f-aca4-4fe02f6d838f" SciMLTesting = "09d9d899-5365-40a9-917a-5f67fddea283" Test = "8dfed614-e22c-5e08-85e1-65c5234f0b40" +blis_jll = "6136c539-28a5-5bf0-87cc-b183200dce32" [sources] LinearSolve = {path = "../.."} [compat] +AllocCheck = "0.2" Aqua = "0.8" JET = "0.9, 0.11" +LAPACK_jll = "3" +LinearAlgebra = "1.10" LinearSolve = "5" SafeTestsets = "0.1, 1" SciMLTesting = "2.1" Test = "<0.0.1, 1" +blis_jll = "0.9.0" julia = "1.10" diff --git a/test/qa/allocations.jl b/test/qa/allocations.jl new file mode 100644 index 000000000..f6ae16ec4 --- /dev/null +++ b/test/qa/allocations.jl @@ -0,0 +1,64 @@ +using AllocCheck, LinearAlgebra, LinearSolve, Test + +if Sys.islinux() + import LAPACK_jll, blis_jll +end + +@check_allocs function allocation_checked_direct_lu_refactor_solve!( + cache, Awork, A, alg + ) + copyto!(Awork, A) + cache.A = Awork + info = LinearSolve._direct_lu_factorize!(cache.cacheval, Awork, alg) + iszero(info) || return info + LinearSolve._direct_lu_solve!(cache.cacheval, cache.u, cache.b, alg) + cache.isfresh = false + return info +end + +function test_allocation_free_refactorization(alg, ::Type{T}) where {T} + A1 = T[4 1; 2 3] + A2 = T[3 -1; 1 2] + b = T[1, 2] + cache = init(LinearProblem(copy(A1), copy(b)), alg) + Awork = cache.A + + @test solve!(cache).u ≈ A1 \ b + info = allocation_checked_direct_lu_refactor_solve!(cache, Awork, A2, alg) + @test iszero(info) + @test cache.u ≈ A2 \ b + + copyto!(Awork, A1) + cache.A = Awork + @test solve!(cache).u ≈ A1 \ b + copyto!(Awork, A2) + cache.A = Awork + if VERSION >= v"1.12" + @test @allocated(solve!(cache)) == 0 + else + solve!(cache) + end + return @test cache.u ≈ A2 \ b +end + +@testset "Direct BLAS refactorization solve! is allocation-free" begin + if LinearSolve.useopenblas + for T in (Float32, Float64, ComplexF32, ComplexF64) + test_allocation_free_refactorization(OpenBLASLUFactorization(), T) + end + end + + if Base.get_extension(LinearSolve, :LinearSolveBLISExt) !== nothing + for T in (Float32, Float64, ComplexF32, ComplexF64) + test_allocation_free_refactorization(LinearSolve.BLISLUFactorization(), T) + end + end +end + +if LinearSolve.appleaccelerate_isavailable() + @testset "Apple Accelerate refactorization solve! is allocation-free" begin + for T in (Float32, Float64, ComplexF32, ComplexF64) + test_allocation_free_refactorization(AppleAccelerateLUFactorization(), T) + end + end +end diff --git a/test/qa/jet.jl b/test/qa/jet.jl index 8aadbf829..fd995c9f6 100644 --- a/test/qa/jet.jl +++ b/test/qa/jet.jl @@ -171,10 +171,9 @@ end @testset "JET Tests for Default Solver" begin # Test the default solver selection - # These tests have various runtime dispatch issues in stdlib code: - # - Dense: Captured variables in appleaccelerate.jl (platform-specific) - # - Sparse: Runtime dispatch in SparseArrays stdlib, Base.show, etc. - JET.@test_opt solve(prob) broken = true + # Julia 1.10 reports runtime dispatch through stdlib and Krylov fallback paths. + JET.@test_opt solve(prob) broken = VERSION < v"1.12.0-" + # Sparse has runtime dispatch in SparseArrays stdlib, Base.show, etc. JET.@test_opt solve(prob_sparse) broken = true end diff --git a/test/runtests.jl b/test/runtests.jl index 5913fc259..0887e70fa 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -66,6 +66,7 @@ else @time @safetestset "Batched RHS" include("Core/batch.jl") @time @safetestset "GESV Factorization" include("Core/gesv.jl") @time @safetestset "LU Refactorization Reuse" include("Core/lu_refactorization.jl") + @time @safetestset "Direct BLAS Refactorization Reuse" include("Core/direct_blas_refactorization.jl") @time @safetestset "Lightweight Solution (no cache)" include("Core/lightweight_solution.jl") @time @safetestset "Return codes" include("Core/retcodes.jl") @time @safetestset "Re-solve" include("Core/resolve.jl") @@ -87,6 +88,12 @@ else return @time @safetestset "SpecializingFactorizations" include("Core/specializing_factorizations.jl") end, groups = Dict( + "AppleAccelerate" => function () + @time @safetestset "Apple Accelerate Refactorization Reuse" include("Core/lu_refactorization.jl") + @time @safetestset "Apple Accelerate Mixed Precision" include("Core/test_mixed_precision.jl") + activate_group_env(joinpath(@__DIR__, "qa")) + return @time @safetestset "Apple Accelerate Allocation QA" include("qa/allocations.jl") + end, # STRUMPACK runs in the base env: STRUMPACK_jll is a base test dep (the # Core suite also probes the STRUMPACK extension), so this group adds no # deps. @@ -219,6 +226,7 @@ else activate_group_env(joinpath(@__DIR__, "qa")) @time @safetestset "Quality Assurance" include("qa/qa.jl") @time @safetestset "JET Tests" include("qa/jet.jl") + @time @safetestset "Allocation QA" include("qa/allocations.jl") end return nothing end, diff --git a/test/test_groups.toml b/test/test_groups.toml index 9a29a1105..e1fa293f4 100644 --- a/test/test_groups.toml +++ b/test/test_groups.toml @@ -23,6 +23,12 @@ versions = ["lts", "1", "pre"] runner = "self-hosted" num_threads = 2 +# Keep the direct Accelerate refactorization allocation contract covered on Apple hardware. +[AppleAccelerate] +versions = ["lts", "1", "pre"] +os = ["macos-latest"] +num_threads = 2 + [LinearSolveSTRUMPACK] versions = ["lts", "1", "pre"] runner = "self-hosted"