diff --git a/test/testsuite/ad_utils.jl b/test/testsuite/ad_utils.jl index def762889..1ab76e2d1 100644 --- a/test/testsuite/ad_utils.jl +++ b/test/testsuite/ad_utils.jl @@ -59,6 +59,27 @@ test in-place Hermitian eigendecomposition rules via Mooncake's non-primitive AD """ eigh!_wrapper(f!, A, alg) = (F = f!(project_hermitian!(A), alg); MatrixAlgebraKit.zero!(A); F) +""" + eig_vals_wrapper(f, A, alg) + +Wrapper that sorts the eigenvalues returned by `f(A, alg)` by modulus and then by imaginary part. +LAPACK's ordering of the eigenvalues can change discontinuously under small perturbations of +`A`, which breaks finite-difference checks. The ordering imposed here is smooth for matrices +built with `make_eig_matrix`, whose eigenvalues have distinct moduli up to conjugate pairs. +""" +eig_vals_wrapper(f, A, alg) = sort_eigvals(f(A, alg)) + +""" + eig_vals!_wrapper(f!, A, alg) + +In-place variant of [`eig_vals_wrapper`](@ref), which zeros `A` after calling `f!`. +""" +eig_vals!_wrapper(f!, A, alg) = sort_eigvals(call_and_zero!(f!, A, alg)) + +# sortperm is used here because Mooncake CAN differentiate that on CUDA, +# but CANNOT differentiate sort +sort_eigvals(D) = D[sortperm(collect(D); by = λ -> (abs(λ), imag(λ)))] + """ qr_gauge_invariant_wrapper(f, A, alg, r) @@ -149,11 +170,27 @@ function stabilize_eigvals!(D::AbstractVector) n = maximum(p) # rescale eigenvalues so that they lie on distinct radii in the complex plane # that are chosen randomly in non-overlapping intervals [10 * k/n, 10 * (k+0.5)/n)] for k=1,...,n - radii = 10 .* ((1:n) .+ rand(real(eltype(D)), n) ./ 2) ./ n + radii = 10 .* ((1:n) .+ rand(rng, real(eltype(D)), n) ./ 2) ./ n hD = sign.(collect(D)) .* radii[p] copyto!(D, hD) return D end +""" + midgap_tol(vals) + +Return a truncation tolerance halfway across the widest gap between consecutive values of +`abs.(vals)`, restricted to the middle half so that truncation keeps a nontrivial subset. +This keeps the number of retained values fixed under the perturbations used by +finite-difference checks. +""" +function midgap_tol(vals) + s = sort!(collect(abs.(vals))) + n = length(s) + gaps = (max(1, n ÷ 4)):(min(n - 1, (3n) ÷ 4)) + _, i = findmax(i -> s[i + 1] - s[i], gaps) + return (s[gaps[i]] + s[gaps[i] + 1]) / 2 +end + function make_eig_matrix(T, sz) A = instantiate_matrix(T, sz) D, V = eig_full(A) diff --git a/test/testsuite/chainrules.jl b/test/testsuite/chainrules.jl index 631d95576..df95e7026 100755 --- a/test/testsuite/chainrules.jl +++ b/test/testsuite/chainrules.jl @@ -442,7 +442,7 @@ function test_chainrules_eigh( @test isequal(ΔDVtrunc, ΔDVtrunc_copy) end D, ΔD = ad_eigh_vals_setup(A / 2) - truncalg = TruncatedAlgorithm(alg, trunctol(; atol = maximum(abs, D) / 2)) + truncalg = TruncatedAlgorithm(alg, trunctol(; atol = midgap_tol(eigh_vals(A)))) DV, DVtrunc, ΔDV, ΔDVtrunc = ad_eigh_trunc_setup(A, truncalg) ind = MatrixAlgebraKit.findtruncated(diagview(DV[1]), truncalg.trunc) ot = (ΔDVtrunc..., zero(real(T))) @@ -601,7 +601,7 @@ function test_chainrules_svd( @test isequal(ΔUSVᴴtrunc, ΔUSVᴴtrunc_copy) end S, ΔS = ad_svd_vals_setup(A) - truncalg = TruncatedAlgorithm(alg, trunctol(atol = S[1, 1] / 2)) + truncalg = TruncatedAlgorithm(alg, trunctol(atol = midgap_tol(S))) USVᴴ, _, ΔUSVᴴ, ΔUSVᴴtrunc = ad_svd_trunc_setup(A, truncalg) ot = (ΔUSVᴴtrunc..., zero(real(T))) ot_copy = deepcopy(ot) @@ -625,7 +625,7 @@ function test_chainrules_svd( dA1 = MatrixAlgebraKit.svd_pullback!(zero(A), A, USVᴴ, ΔUSVᴴtrunc, ind) dA2 = MatrixAlgebraKit.svd_trunc_pullback!(zero(A), A, (Utrunc, Strunc, Vᴴtrunc), ΔUSVᴴtrunc) @test isapprox(dA1, dA2; atol = atol, rtol = rtol) - trunc = trunctol(; atol = S[1, 1] / 2) + trunc = truncalg.trunc ind = MatrixAlgebraKit.findtruncated(diagview(S), trunc) ot = (ΔUSVᴴtrunc..., zero(real(T))) ot_copy = deepcopy(ot) diff --git a/test/testsuite/enzyme/eig.jl b/test/testsuite/enzyme/eig.jl index 4eaf878f3..b9e609d81 100644 --- a/test/testsuite/enzyme/eig.jl +++ b/test/testsuite/enzyme/eig.jl @@ -19,7 +19,7 @@ Test the Enzyme foward- and reverse-mode AD rule for `eig_full` and its in-place """ function test_enzyme_eig_full( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "eig_full: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -40,17 +40,17 @@ Test the Enzyme forward- and reverse-mode AD rule for `eig_vals` and its in-plac """ function test_enzyme_eig_vals( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "eig_vals: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) A = make_eig_matrix(T, sz) alg = MatrixAlgebraKit.select_algorithm(eig_vals, A) D, ΔD = ad_eig_vals_setup(A) - test_reverse(eig_vals, RT, (A, TA), (alg, Const); atol, rtol, output_tangent = ΔD, fdm) - test_reverse(call_and_zero!, RT, (eig_vals!, Const), (copy(A), TA), (alg, Const); atol, rtol, output_tangent = ΔD, fdm) - test_forward(eig_vals, RT, (A, TA), (alg, Const); atol, rtol, fdm) - test_forward(call_and_zero!, RT, (eig_vals!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm) + test_reverse(eig_vals_wrapper, RT, (eig_vals, Const), (A, TA), (alg, Const); atol, rtol, output_tangent = ΔD, fdm) + test_reverse(eig_vals!_wrapper, RT, (eig_vals!, Const), (copy(A), TA), (alg, Const); atol, rtol, output_tangent = ΔD, fdm) + test_forward(eig_vals_wrapper, RT, (eig_vals, Const), (A, TA), (alg, Const); atol, rtol, fdm) + test_forward(eig_vals!_wrapper, RT, (eig_vals!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm) end end @@ -62,7 +62,7 @@ in-place variants, over a range of truncation ranks and a tolerance-based trunca """ function test_enzyme_eig_trunc( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "eig_trunc reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -83,7 +83,7 @@ function test_enzyme_eig_trunc( @testset "trunctol" begin A = make_eig_matrix(T, sz) D = eig_vals(A) - trunc = trunctol(atol = maximum(abs, D) / 2; by = abs) + trunc = trunctol(atol = midgap_tol(D); by = abs) truncalg = TruncatedAlgorithm(alg, trunc) DV, _, ΔDV, ΔDVtrunc = ad_eig_trunc_setup(A, truncalg) test_reverse(eig_trunc_no_error, RT, (A, TA), (truncalg, Const); atol, rtol, output_tangent = ΔDVtrunc, fdm) diff --git a/test/testsuite/enzyme/eigh.jl b/test/testsuite/enzyme/eigh.jl index 8143678ec..84cf41842 100644 --- a/test/testsuite/enzyme/eigh.jl +++ b/test/testsuite/enzyme/eigh.jl @@ -19,7 +19,7 @@ Test the Enzyme forward- and reverse-mode AD rule for `eigh_full` and its in-pla """ function test_enzyme_eigh_full( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "eigh_full: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -40,7 +40,7 @@ Test the Enzyme forward- and reverse-mode AD rule for `eigh_vals` and its in-pla """ function test_enzyme_eigh_vals( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "eigh_vals: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -62,7 +62,7 @@ in-place variants, over a range of truncation ranks. """ function test_enzyme_eigh_trunc( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "eigh_trunc reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -82,8 +82,8 @@ function test_enzyme_eigh_trunc( end @testset "trunctol" begin A = make_eigh_matrix(T, sz) - D = eigh_vals(A / 2, alg) - trunc = trunctol(; atol = maximum(abs, D) / 2) + D = eigh_vals(A, alg) + trunc = trunctol(; atol = midgap_tol(D)) truncalg = TruncatedAlgorithm(alg, trunc) DV, _, ΔDV, ΔDVtrunc = ad_eigh_trunc_setup(A, truncalg) test_reverse(eigh_wrapper, RT, (eigh_trunc_no_error, Const), (A, TA), (truncalg, Const); atol, rtol, output_tangent = ΔDVtrunc, fdm) diff --git a/test/testsuite/enzyme/lq.jl b/test/testsuite/enzyme/lq.jl index 9b56ce110..2e35b3d56 100644 --- a/test/testsuite/enzyme/lq.jl +++ b/test/testsuite/enzyme/lq.jl @@ -15,7 +15,7 @@ end function test_enzyme_lq_compact( T::Type, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "lq_compact: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -31,7 +31,7 @@ end function test_enzyme_lq_compact_rank_deficient( T::Type, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "lq_compact rank deficient A: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -52,7 +52,7 @@ end function test_enzyme_lq_full( T::Type, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "lq_full reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -70,7 +70,7 @@ end function test_enzyme_lq_null( T::Type, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "lq_null reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) diff --git a/test/testsuite/enzyme/orthnull.jl b/test/testsuite/enzyme/orthnull.jl index 05a1640ea..dec6969f1 100644 --- a/test/testsuite/enzyme/orthnull.jl +++ b/test/testsuite/enzyme/orthnull.jl @@ -22,7 +22,7 @@ algorithms, and their in-place variants. """ function test_enzyme_left_orth( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "left_orth reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -61,7 +61,7 @@ algorithms, and their in-place variants. """ function test_enzyme_right_orth( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "right_orth reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -99,7 +99,7 @@ in-place variant. """ function test_enzyme_left_null( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "left_null: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -124,7 +124,7 @@ in-place variant. """ function test_enzyme_right_null( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "right_null: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) diff --git a/test/testsuite/enzyme/polar.jl b/test/testsuite/enzyme/polar.jl index e342f416d..21cf54f67 100644 --- a/test/testsuite/enzyme/polar.jl +++ b/test/testsuite/enzyme/polar.jl @@ -19,7 +19,7 @@ Only runs for tall or square matrices (`m >= n`). """ function test_enzyme_left_polar( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "left_polar: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) A = instantiate_matrix(T, sz) @@ -44,7 +44,7 @@ Only runs for wide or square matrices (`m <= n`). """ function test_enzyme_right_polar( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "right_polar: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) A = instantiate_matrix(T, sz) diff --git a/test/testsuite/enzyme/projections.jl b/test/testsuite/enzyme/projections.jl index 2c5240a2f..c40c0006c 100644 --- a/test/testsuite/enzyme/projections.jl +++ b/test/testsuite/enzyme/projections.jl @@ -19,7 +19,7 @@ Test the Enzyme forward- and reverse-mode AD rule for `project_hermitian` and it """ function test_enzyme_project_hermitian( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "project_hermitian: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -42,7 +42,7 @@ Test the Enzyme forward- and reverse-mode AD rule for `project_antihermitian` an """ function test_enzyme_project_antihermitian( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "project_antihermitian: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) diff --git a/test/testsuite/enzyme/qr.jl b/test/testsuite/enzyme/qr.jl index f3b8f0901..f7caaa2f5 100644 --- a/test/testsuite/enzyme/qr.jl +++ b/test/testsuite/enzyme/qr.jl @@ -15,7 +15,7 @@ end function test_enzyme_qr_compact( T::Type, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "qr_compact reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -31,7 +31,7 @@ end function test_enzyme_qr_compact_rank_deficient( T::Type, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "qr_compact rank deficient A reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -52,7 +52,7 @@ end function test_enzyme_qr_full( T::Type, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "qr_full reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -70,7 +70,7 @@ end function test_enzyme_qr_null( T::Type, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "qr_null reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) diff --git a/test/testsuite/enzyme/svd.jl b/test/testsuite/enzyme/svd.jl index cddd1b8ee..9155021f1 100644 --- a/test/testsuite/enzyme/svd.jl +++ b/test/testsuite/enzyme/svd.jl @@ -15,7 +15,7 @@ Test the Enzyme forward- and reverse-mode AD rule for `svd_compact` and its in-p """ function test_enzyme_svd_compact( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "svd_compact: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -37,7 +37,7 @@ gauge-dependent extra columns of `U` and rows of `Vᴴ` are zeroed out in the co """ function test_enzyme_svd_full( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "svd_full: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -60,7 +60,7 @@ Test the Enzyme forward- and reverse-mode AD rule for `svd_vals` and its in-plac """ function test_enzyme_svd_vals( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "svd_vals: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -82,7 +82,7 @@ in-place variants, over a range of truncation ranks and a tolerance-based trunca """ function test_enzyme_svd_trunc( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) return @testset "svd_trunc reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) @@ -101,7 +101,7 @@ function test_enzyme_svd_trunc( end @testset "trunctol" begin S = svd_vals(A, alg) - trunc = trunctol(atol = maximum(S) / 2) + trunc = trunctol(atol = midgap_tol(S)) truncalg = TruncatedAlgorithm(alg, trunc) USVᴴ, _, ΔUSVᴴ, ΔUSVᴴtrunc = ad_svd_trunc_setup(A, truncalg) test_reverse(svd_trunc_no_error, RT, (A, TA), (truncalg, Const); atol, rtol, output_tangent = ΔUSVᴴtrunc, fdm) diff --git a/test/testsuite/mooncake/eig.jl b/test/testsuite/mooncake/eig.jl index 8c9c29199..51a1049f8 100644 --- a/test/testsuite/mooncake/eig.jl +++ b/test/testsuite/mooncake/eig.jl @@ -19,7 +19,7 @@ Test the Mooncake forward- and reverse-mode AD rule for `eig_full` and its in-pl """ function test_mooncake_eig_full( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "eig_full" begin A = make_eig_matrix(T, sz) @@ -65,7 +65,7 @@ Test the Mooncake forward- and reverse-mode AD rule for `eig_vals` and its in-pl """ function test_mooncake_eig_vals( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "eig_vals" begin A = make_eig_matrix(T, sz) @@ -74,8 +74,8 @@ function test_mooncake_eig_vals( output_tangent = Mooncake.randn_tangent(rng, D) Mooncake.TestUtils.test_rule( - rng, eig_vals, A, alg; - output_tangent, atol, rtol + rng, eig_vals_wrapper, eig_vals, A, alg; + output_tangent, atol, rtol, is_primitive = false ) if A isa Diagonal{<:Complex} A2 = copy(A) @@ -90,7 +90,7 @@ function test_mooncake_eig_vals( ) end Mooncake.TestUtils.test_rule( - rng, call_and_zero!, eig_vals!, A, alg; + rng, eig_vals!_wrapper, eig_vals!, A, alg; output_tangent, atol, rtol, is_primitive = false ) end @@ -104,7 +104,7 @@ in-place variants, over a range of truncation ranks and a tolerance-based trunca """ function test_mooncake_eig_trunc( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "eig_trunc" begin A = make_eig_matrix(T, sz) @@ -144,7 +144,7 @@ function test_mooncake_eig_trunc( @testset "trunctol" begin D = eig_vals(A) - trunc = trunctol(atol = maximum(abs, D) / 2; by = abs) + trunc = trunctol(atol = midgap_tol(D); by = abs) alg_trunc = TruncatedAlgorithm(alg, trunc) DV, DVtrunc, ΔDV_arrays, ΔDVtrunc_arrays = ad_eig_trunc_setup(A, alg_trunc) diff --git a/test/testsuite/mooncake/eigh.jl b/test/testsuite/mooncake/eigh.jl index 481c4ea54..81c5bd16b 100644 --- a/test/testsuite/mooncake/eigh.jl +++ b/test/testsuite/mooncake/eigh.jl @@ -19,7 +19,7 @@ Test the Mooncake reverse-mode AD rule for `eigh_full` and its in-place variant. """ function test_mooncake_eigh_full( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "eigh_full" begin A = make_eigh_matrix(T, sz) @@ -55,7 +55,7 @@ Test the Mooncake reverse-mode AD rule for `eigh_vals` and its in-place variant. """ function test_mooncake_eigh_vals( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "eigh_vals" begin A = make_eigh_matrix(T, sz) @@ -94,7 +94,7 @@ in-place variants, over a range of truncation ranks and a tolerance-based trunca """ function test_mooncake_eigh_trunc( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "eigh_trunc" begin A = make_eigh_matrix(T, sz) @@ -134,7 +134,7 @@ function test_mooncake_eigh_trunc( @testset "trunctol" begin D = eigh_vals(A) - trunc = trunctol(atol = maximum(abs, D) / 2; by = abs) + trunc = trunctol(atol = midgap_tol(D); by = abs) alg_trunc = TruncatedAlgorithm(alg, trunc) DV, DVtrunc, ΔDV_arrays, ΔDVtrunc_arrays = ad_eigh_trunc_setup(A, alg_trunc) diff --git a/test/testsuite/mooncake/lq.jl b/test/testsuite/mooncake/lq.jl index f947095fc..3de9dac4c 100644 --- a/test/testsuite/mooncake/lq.jl +++ b/test/testsuite/mooncake/lq.jl @@ -19,7 +19,7 @@ Test the Mooncake forward- and reverse-mode AD rule for `lq_compact` and its in- """ function test_mooncake_lq_compact( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "lq_compact" begin A = instantiate_matrix(T, sz) @@ -71,7 +71,7 @@ Test the Mooncake forward- and reverse-mode AD rule for `lq_full` and its in-pla """ function test_mooncake_lq_full( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "lq_full" begin A = instantiate_matrix(T, sz) @@ -107,7 +107,7 @@ Test the Mooncake forward- and reverse-mode AD rule for `lq_null` and its in-pla """ function test_mooncake_lq_null( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "lq_null" begin A = instantiate_matrix(T, sz) diff --git a/test/testsuite/mooncake/orthnull.jl b/test/testsuite/mooncake/orthnull.jl index f0285634a..d3bc5b84f 100644 --- a/test/testsuite/mooncake/orthnull.jl +++ b/test/testsuite/mooncake/orthnull.jl @@ -22,7 +22,7 @@ algorithms, and their in-place variants. """ function test_mooncake_left_orth( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "left_orth" begin A = instantiate_matrix(T, sz) @@ -70,7 +70,7 @@ algorithms, and their in-place variants. """ function test_mooncake_right_orth( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "right_orth" begin A = instantiate_matrix(T, sz) @@ -118,7 +118,7 @@ in-place variant. """ function test_mooncake_left_null( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "left_null" begin A = instantiate_matrix(T, sz) @@ -157,7 +157,7 @@ in-place variant. """ function test_mooncake_right_null( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "right_null" begin A = instantiate_matrix(T, sz) diff --git a/test/testsuite/mooncake/polar.jl b/test/testsuite/mooncake/polar.jl index fee32e71e..5170b5d57 100644 --- a/test/testsuite/mooncake/polar.jl +++ b/test/testsuite/mooncake/polar.jl @@ -19,7 +19,7 @@ Only runs for tall or square matrices (`m >= n`). """ function test_mooncake_left_polar( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "left_polar" begin A = instantiate_matrix(T, sz) @@ -48,7 +48,7 @@ Only runs for wide or square matrices (`m <= n`). """ function test_mooncake_right_polar( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "right_polar" begin A = instantiate_matrix(T, sz) diff --git a/test/testsuite/mooncake/projections.jl b/test/testsuite/mooncake/projections.jl index 359c9cbe9..f6baf91ec 100644 --- a/test/testsuite/mooncake/projections.jl +++ b/test/testsuite/mooncake/projections.jl @@ -19,7 +19,7 @@ Test the Mooncake forward- and reverse-mode AD rule for `project_hermitian` and """ function test_mooncake_project_hermitian( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "project_hermitian" begin A = instantiate_matrix(T, sz) @@ -47,7 +47,7 @@ Test the Mooncake forward- and reverse-mode AD rule for `project_antihermitian` """ function test_mooncake_project_antihermitian( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "project_antihermitian" begin A = instantiate_matrix(T, sz) diff --git a/test/testsuite/mooncake/qr.jl b/test/testsuite/mooncake/qr.jl index 830e3472b..96a35e309 100644 --- a/test/testsuite/mooncake/qr.jl +++ b/test/testsuite/mooncake/qr.jl @@ -19,7 +19,7 @@ Test the Mooncake forward- and reverse-mode AD rule for `qr_compact` and its in- """ function test_mooncake_qr_compact( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "qr_compact" begin A = instantiate_matrix(T, sz) @@ -71,7 +71,7 @@ Test the Mooncake forward- and reverse-mode AD rule for `qr_full` and its in-pla """ function test_mooncake_qr_full( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "qr_full" begin A = instantiate_matrix(T, sz) @@ -107,7 +107,7 @@ Test the Mooncake forward- and reverse-mode AD rule for `qr_null` and its in-pla """ function test_mooncake_qr_null( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "qr_null" begin A = instantiate_matrix(T, sz) diff --git a/test/testsuite/mooncake/svd.jl b/test/testsuite/mooncake/svd.jl index de9dbc543..f13faa874 100644 --- a/test/testsuite/mooncake/svd.jl +++ b/test/testsuite/mooncake/svd.jl @@ -20,7 +20,7 @@ Test the Mooncake forward- and reverse-mode AD rule for `svd_compact` and its in """ function test_mooncake_svd_compact( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "svd_compact" begin A = instantiate_matrix(T, sz) @@ -47,7 +47,7 @@ gauge-dependent extra columns of `U` and rows of `Vᴴ` are zeroed out in the co """ function test_mooncake_svd_full( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "svd_full" begin A = instantiate_matrix(T, sz) @@ -75,7 +75,7 @@ Test the Mooncake forward- and reverse-mode AD rule for `svd_vals` and its in-pl """ function test_mooncake_svd_vals( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "svd_vals" begin A = instantiate_matrix(T, sz) @@ -102,7 +102,7 @@ in-place variants, over a range of truncation ranks and a tolerance-based trunca """ function test_mooncake_svd_trunc( T, sz; - rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T) + rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T) ) return @testset "svd_trunc" begin A = instantiate_matrix(T, sz) @@ -143,7 +143,7 @@ function test_mooncake_svd_trunc( @testset "trunctol" begin S = svd_vals(A) - trunc = trunctol(atol = maximum(S) / 2) + trunc = trunctol(atol = midgap_tol(S)) alg_trunc = TruncatedAlgorithm(alg, trunc) USVᴴ, USVᴴtrunc, ΔUSVᴴ_arrays, ΔUSVᴴtrunc_arrays = ad_svd_trunc_setup(A, alg_trunc)