From 8fab5a51e2cbe35478cf6a929241b45955ca94b4 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 24 Sep 2026 16:11:43 +0200 Subject: [PATCH 01/13] Forward rules for QR/LQ --- .../MatrixAlgebraKitEnzymeExt.jl | 38 ++++++++- .../MatrixAlgebraKitMooncakeExt.jl | 34 ++++++++ src/MatrixAlgebraKit.jl | 2 + src/pushforwards/lq.jl | 76 ++++++++++++++++++ src/pushforwards/qr.jl | 75 ++++++++++++++++++ test/testsuite/ad_utils.jl | 78 +++++++++++++++++++ test/testsuite/enzyme/lq.jl | 25 ++++-- test/testsuite/enzyme/orthnull.jl | 34 +++++--- test/testsuite/enzyme/qr.jl | 21 ++++- test/testsuite/mooncake/lq.jl | 41 ++++++++-- test/testsuite/mooncake/orthnull.jl | 30 +++++-- test/testsuite/mooncake/qr.jl | 41 ++++++++-- 12 files changed, 454 insertions(+), 41 deletions(-) create mode 100644 src/pushforwards/lq.jl create mode 100644 src/pushforwards/qr.jl diff --git a/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl b/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl index b252b82a1..90f93883b 100644 --- a/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl +++ b/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl @@ -4,7 +4,9 @@ using MatrixAlgebraKit using MatrixAlgebraKit: copy_input, initialize_output, zero!, has_equal_storage using MatrixAlgebraKit: diagview, inv_safe, truncate using MatrixAlgebraKit: qr_pullback!, lq_pullback! +using MatrixAlgebraKit: qr_pushforward!, lq_pushforward! using MatrixAlgebraKit: qr_null_pullback!, lq_null_pullback! +using MatrixAlgebraKit: qr_null_pushforward!, lq_null_pushforward! using MatrixAlgebraKit: eig_pullback!, eigh_pullback!, eig_vals_pullback!, eigh_vals_pullback! using MatrixAlgebraKit: eig_pushforward!, eigh_pushforward!, eig_vals_pushforward!, eigh_vals_pushforward! using MatrixAlgebraKit: svd_pullback!, svd_vals_pullback! @@ -121,6 +123,10 @@ for (f, pf) in ( (:left_polar!, :left_polar_pushforward!), (:eigh_full!, :eigh_pushforward!), (:eig_full!, :eig_pushforward!), + (:qr_full!, :qr_pushforward!), + (:lq_full!, :lq_pushforward!), + (:qr_compact!, :qr_pushforward!), + (:lq_compact!, :lq_pushforward!), ) @eval begin function EnzymeRules.forward( @@ -152,9 +158,9 @@ for (f, pf) in ( end end -for (f, pb) in ( - (qr_null!, qr_null_pullback!), - (lq_null!, lq_null_pullback!), +for (f, pb, pf) in ( + (qr_null!, qr_null_pullback!, qr_null_pushforward!), + (lq_null!, lq_null_pullback!, lq_null_pushforward!), ) @eval begin function EnzymeRules.augmented_primal( @@ -203,6 +209,32 @@ for (f, pb) in ( !isa(arg, Const) && make_zero!(arg.dval) return (nothing, nothing, nothing) end + function EnzymeRules.forward( + config::EnzymeRules.FwdConfigWidth{1}, + func::Const{typeof($f)}, + ::Type{RT}, + A::Annotation, + arg::Annotation, + alg::Const{<:MatrixAlgebraKit.AbstractAlgorithm}, + ) where {RT} + # here, A IS directly used in the pushforward, and overwritten + # in the primal call, so we MUST copy its value + Ac = copy(A.val) + $f(A.val, arg.val, alg.val) + if !isa(A, Const) && !isa(arg, Const) + $pf(A.dval, Ac, arg.val, arg.dval) + end + !isa(A, Const) && make_zero!(A.dval) + if EnzymeRules.needs_primal(config) && EnzymeRules.needs_shadow(config) + return arg + elseif EnzymeRules.needs_primal(config) + return arg.val + elseif EnzymeRules.needs_shadow(config) + return arg.dval + else + return nothing + end + end end end diff --git a/ext/MatrixAlgebraKitMooncakeExt/MatrixAlgebraKitMooncakeExt.jl b/ext/MatrixAlgebraKitMooncakeExt/MatrixAlgebraKitMooncakeExt.jl index 4d601a720..c3e7f4b1b 100644 --- a/ext/MatrixAlgebraKitMooncakeExt/MatrixAlgebraKitMooncakeExt.jl +++ b/ext/MatrixAlgebraKitMooncakeExt/MatrixAlgebraKitMooncakeExt.jl @@ -6,6 +6,8 @@ using MatrixAlgebraKit using MatrixAlgebraKit: inv_safe, diagview, copy_input, initialize_output, zero!, has_equal_storage using MatrixAlgebraKit: qr_pullback!, lq_pullback! using MatrixAlgebraKit: qr_null_pullback!, lq_null_pullback! +using MatrixAlgebraKit: qr_pushforward!, lq_pushforward! +using MatrixAlgebraKit: qr_null_pushforward!, lq_null_pushforward! using MatrixAlgebraKit: eig_pullback!, eigh_pullback!, eig_vals_pullback! using MatrixAlgebraKit: eig_pushforward!, eig_vals_pushforward! using MatrixAlgebraKit: eigh_pushforward!, eigh_vals_pushforward! @@ -112,6 +114,10 @@ end for (f!, f, pf) in ( (:left_polar!, :left_polar, :left_polar_pushforward!), (:right_polar!, :right_polar, :right_polar_pushforward!), + (:qr_full!, :qr_full, :qr_pushforward!), + (:qr_compact!, :qr_compact, :qr_pushforward!), + (:lq_full!, :lq_full, :lq_pushforward!), + (:lq_compact!, :lq_compact, :lq_pushforward!), (:eig_full!, :eig_full, :eig_pushforward!), (:eigh_full!, :eigh_full, :eigh_pushforward!), ) @@ -180,6 +186,34 @@ for (f!, f, pb, adj) in ( end end +for (f!, f, pf) in ( + (:qr_null!, :qr_null, :qr_null_pushforward!), + (:lq_null!, :lq_null, :lq_null_pushforward!), + ) + @eval begin + @is_primitive Mooncake.DefaultCtx Mooncake.ForwardMode Tuple{typeof($f!), Any, Any, MatrixAlgebraKit.AbstractAlgorithm} + function Mooncake.frule!!(f_df::Dual{typeof($f!)}, A_dA::Dual, arg_darg::Dual, alg_dalg::Dual{<:MatrixAlgebraKit.AbstractAlgorithm}) + A, dA = arrayify(A_dA) + arg, darg = arrayify(arg_darg) + # the pushforward needs the original A, which is destroyed by $f! + Ac = copy(A) + $f!(A, arg, Mooncake.primal(alg_dalg)) + $pf(dA, Ac, arg, darg) + return arg_darg + end + @is_primitive Mooncake.DefaultCtx Mooncake.ForwardMode Tuple{typeof($f), Any, MatrixAlgebraKit.AbstractAlgorithm} + function Mooncake.frule!!(f_df::Dual{typeof($f)}, A_dA::Dual, alg_dalg::Dual{<:MatrixAlgebraKit.AbstractAlgorithm}) + A, dA = arrayify(A_dA) + output = $f(A, Mooncake.primal(alg_dalg)) + doutput = Mooncake.zero_tangent(output) + output_dual = Dual(output, doutput) + arg, darg = arrayify(output_dual) + $pf(dA, A, arg, darg) + return output_dual + end + end +end + for (f!, f, f_full, f_full!, pb, pf, adj) in ( (:eig_vals!, :eig_vals, :eig_full, :eig_full!, :eig_vals_pullback!, :eig_vals_pushforward!, :eig_vals_adjoint), (:eigh_vals!, :eigh_vals, :eigh_full, :eigh_full!, :eigh_vals_pullback!, :eigh_vals_pushforward!, :eigh_vals_adjoint), diff --git a/src/MatrixAlgebraKit.jl b/src/MatrixAlgebraKit.jl index a3e007df2..3d9662137 100644 --- a/src/MatrixAlgebraKit.jl +++ b/src/MatrixAlgebraKit.jl @@ -142,6 +142,8 @@ include("pushforwards/polar.jl") include("pushforwards/eig.jl") include("pushforwards/eigh.jl") include("pushforwards/svd.jl") +include("pushforwards/qr.jl") +include("pushforwards/lq.jl") include("precompile.jl") diff --git a/src/pushforwards/lq.jl b/src/pushforwards/lq.jl new file mode 100644 index 000000000..80d1a59a5 --- /dev/null +++ b/src/pushforwards/lq.jl @@ -0,0 +1,76 @@ +""" + lq_pushforward!( + ΔA, A, LQ, ΔLQ; + rank_atol::Real = default_pullback_rank_atol(LQ[1]) + ) + +Compute the pushforward `ΔLQ` of the LQ decomposition `LQ` of `lq_compact(A; +positive = true)` or `lq_full(A; positive = true)` given the tangent `ΔA` of `A`. + +If the original matrix `A` is rank-deficient (rank `r < min(size(A)...)`), only the first `r` +rows of `Q` and the first `r` columns of `L` are differentiable, and the tangents of the +remaining rows of `Q` and columns of `L` are set to zero. Similarly, for `lq_full` the extra +rows of `Q` are only determined up to a unitary rotation, and only their gauge-invariant +tangent component along the first `r` rows of `Q` is computed. + +See also [`qr_pushforward!`](@ref). +""" +function lq_pushforward!( + ΔA, A, LQ, ΔLQ; + rank_atol::Real = default_pullback_rank_atol(LQ[1]), kwargs... + ) + L, Q = LQ + ΔL, ΔQ = ΔLQ + m = size(L, 1) + n = size(Q, 2) + minmn = min(m, n) + p = lq_rank(L; rank_atol) + (m, n) == size(ΔA) || throw(DimensionMismatch("size of ΔA ($(size(ΔA))) does not match size of L*Q ($m, $n)")) + + Q₁ = view(Q, 1:p, :) + L₁₁ = LowerTriangular(view(L, 1:p, 1:p)) + L₂₁ = view(L, (p + 1):m, 1:p) + + ΔA₁ = view(ΔA, 1:p, :) + ΔA₂ = view(ΔA, (p + 1):m, :) + + ΔQ₁ = L₁₁ \ ΔA₁ + ΔQ₁Q₁ᴴ = ΔQ₁ * Q₁' + M = ΔQ₁Q₁ᴴ + ΔQ₁Q₁ᴴ' + diagview(M) ./= 2 + view(M, uppertriangularind(M)) .= zero(eltype(M)) + ΔL₁₁ = L₁₁ * M + ΔQ₁ = mul!(ΔQ₁, M, Q₁, -1, 1) + ΔL₂₁ = ΔA₂ * Q₁' + ΔL₂₁ = mul!(ΔL₂₁, L₂₁, Q₁ * ΔQ₁', 1, 1) + + zero!(ΔL) + zero!(ΔQ) + view(ΔQ, 1:p, :) .= ΔQ₁ + view(ΔL, 1:p, 1:p) .= ΔL₁₁ + view(ΔL, (p + 1):m, 1:p) .= ΔL₂₁ + if p == minmn && size(Q, 1) > minmn + Q₃ = view(Q, (minmn + 1):size(Q, 1), :) + ΔQ₃ = view(ΔQ, (minmn + 1):size(Q, 1), :) + mul!(ΔQ₃, Q₃ * ΔQ₁', Q₁, -1, 0) + end + return ΔL, ΔQ +end + +""" + lq_null_pushforward!(ΔA, A, Nᴴ, ΔNᴴ; kwargs...) + +Compute the pushforward `ΔNᴴ` of the left nullspace basis `Nᴴ` of `lq_null(A)` +given the tangent `ΔA` of `A`. + +See also [`lq_pushforward!`](@ref). +""" +function lq_null_pushforward!(ΔA, A, Nᴴ, ΔNᴴ; kwargs...) + if size(Nᴴ, 1) == 0 + zero!(ΔNᴴ) + return ΔNᴴ + end + L, Q = lq_compact(A; positive = true) + X = ldiv!(LowerTriangular(L), ΔA * Nᴴ') + return mul!(ΔNᴴ, X', Q, -1, 0) +end diff --git a/src/pushforwards/qr.jl b/src/pushforwards/qr.jl new file mode 100644 index 000000000..67631d9c6 --- /dev/null +++ b/src/pushforwards/qr.jl @@ -0,0 +1,75 @@ +""" + qr_pushforward!( + ΔA, A, QR, ΔQR; + rank_atol::Real = default_pullback_rank_atol(QR[2]) + ) + +Computes the pushforward `ΔQR` of the QR decomposition `QR` of `qr_compact(A; +positive = true)` or `qr_full(A; positive = true)` given the tangent `ΔA` of `A`. + +If the original matrix `A` is rank-deficient (rank `r < min(size(A)...)`), only the first `r` +columns of `Q` and the first `r` rows of `R` are differentiable, and the tangents of the +remaining columns of `Q` and rows of `R` are set to zero. Similarly, for `qr_full` the extra +columns of `Q` are only determined up to a unitary rotation, and only their gauge-invariant +tangent component along the first `r` columns of `Q` is computed. +""" +function qr_pushforward!( + ΔA, A, QR, ΔQR; + rank_atol::Real = default_pullback_rank_atol(QR[2]), kwargs... + ) + Q, R = QR + ΔQ, ΔR = ΔQR + m = size(Q, 1) + n = size(R, 2) + minmn = min(m, n) + p = qr_rank(R; rank_atol) + (m, n) == size(ΔA) || throw(DimensionMismatch("size of ΔA ($(size(ΔA))) does not match size of Q*R ($m, $n)")) + + Q₁ = view(Q, :, 1:p) + R₁₁ = UpperTriangular(view(R, 1:p, 1:p)) + R₁₂ = view(R, 1:p, (p + 1):n) + + ΔA₁ = view(ΔA, :, 1:p) + ΔA₂ = view(ΔA, :, (p + 1):n) + + ΔQ₁ = ΔA₁ / R₁₁ + Q₁ᴴΔQ₁ = Q₁' * ΔQ₁ + M = Q₁ᴴΔQ₁ + Q₁ᴴΔQ₁' + diagview(M) ./= 2 + view(M, lowertriangularind(M)) .= zero(eltype(M)) + ΔR₁₁ = M * R₁₁ + ΔQ₁ = mul!(ΔQ₁, Q₁, M, -1, 1) + ΔR₁₂ = Q₁' * ΔA₂ + ΔR₁₂ = mul!(ΔR₁₂, ΔQ₁' * Q₁, R₁₂, 1, 1) + + zero!(ΔQ) + zero!(ΔR) + view(ΔQ, :, 1:p) .= ΔQ₁ + view(ΔR, 1:p, 1:p) .= ΔR₁₁ + view(ΔR, 1:p, (p + 1):n) .= ΔR₁₂ + if p == minmn && size(Q, 2) > minmn + # extra columns in the case of qr_full, orthogonality to Q₁ fixes their component along Q₁ + Q₃ = view(Q, :, (minmn + 1):size(Q, 2)) + ΔQ₃ = view(ΔQ, :, (minmn + 1):size(Q, 2)) + mul!(ΔQ₃, Q₁, ΔQ₁' * Q₃, -1, 0) + end + return ΔQ, ΔR +end + +""" + qr_null_pushforward!(ΔA, A, N, ΔN; kwargs...) + +Compute the pushforward `ΔN` of the nullspace basis `N` of `qr_null(A)` given the +tangent `ΔA` of `A`. + +See also [`qr_pushforward!`](@ref). +""" +function qr_null_pushforward!(ΔA, A, N, ΔN; kwargs...) + if size(N, 2) == 0 + zero!(ΔN) + return ΔN + end + Q, R = qr_compact(A; positive = true) + X = ldiv!(UpperTriangular(R)', ΔA' * N) + return mul!(ΔN, Q, X, -1, 0) +end diff --git a/test/testsuite/ad_utils.jl b/test/testsuite/ad_utils.jl index 09558e0ed..def762889 100644 --- a/test/testsuite/ad_utils.jl +++ b/test/testsuite/ad_utils.jl @@ -59,6 +59,84 @@ 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) +""" + qr_gauge_invariant_wrapper(f, A, alg, r) + +Wrapper that calls `Q, R = f(A, alg)` and returns only the parts of the decomposition that +are differentiable functions of `A` if `A` has rank `r`: the first `r` columns of `Q`, the +first `r` rows of `R`, and the projector onto the extra columns of `Q` (for `qr_full`). +Used to test forward-mode QR rules with finite differences, which cannot be restricted to +the gauge-invariant subspace through an `output_tangent`. +""" +qr_gauge_invariant_wrapper(f, A, alg, r) = qr_gauge_invariant_part(f(A, alg), r) + +""" + qr!_gauge_invariant_wrapper(f!, A, alg, r) + +In-place variant of [`qr_gauge_invariant_wrapper`](@ref), which zeros `A` after calling `f!`. +""" +qr!_gauge_invariant_wrapper(f!, A, alg, r) = qr_gauge_invariant_part(call_and_zero!(f!, A, alg), r) + +function qr_gauge_invariant_part((Q, R), r) + minmn = min(size(Q, 1), size(R, 2)) + Q₃ = Q[:, (minmn + 1):end] + return Q[:, 1:r], R[1:r, :], Q₃ * Q₃' +end + +""" + qr_null_gauge_invariant_wrapper(f, A, alg) + +Wrapper that calls `N = f(A, alg)` and returns the projector `N * N'` onto the nullspace, +which, unlike `N` itself, is a differentiable function of `A`. +""" +qr_null_gauge_invariant_wrapper(f, A, alg) = (N = f(A, alg); N * N') + +""" + qr_null!_gauge_invariant_wrapper(f!, A, alg) + +In-place variant of [`qr_null_gauge_invariant_wrapper`](@ref), which zeros `A` after calling `f!`. +""" +qr_null!_gauge_invariant_wrapper(f!, A, alg) = (N = call_and_zero!(f!, A, alg); N * N') + +""" + lq_gauge_invariant_wrapper(f, A, alg, r) + +Wrapper that calls `L, Q = f(A, alg)` and returns only the parts of the decomposition that +are differentiable functions of `A` if `A` has rank `r`: the first `r` columns of `L`, the +first `r` rows of `Q`, and the projector onto the extra rows of `Q` (for `lq_full`). +Used to test forward-mode LQ rules with finite differences, which cannot be restricted to +the gauge-invariant subspace through an `output_tangent`. +""" +lq_gauge_invariant_wrapper(f, A, alg, r) = lq_gauge_invariant_part(f(A, alg), r) + +""" + lq!_gauge_invariant_wrapper(f!, A, alg, r) + +In-place variant of [`lq_gauge_invariant_wrapper`](@ref), which zeros `A` after calling `f!`. +""" +lq!_gauge_invariant_wrapper(f!, A, alg, r) = lq_gauge_invariant_part(call_and_zero!(f!, A, alg), r) + +function lq_gauge_invariant_part((L, Q), r) + minmn = min(size(L, 1), size(Q, 2)) + Q₃ = Q[(minmn + 1):end, :] + return L[:, 1:r], Q[1:r, :], Q₃' * Q₃ +end + +""" + lq_null_gauge_invariant_wrapper(f, A, alg) + +Wrapper that calls `Nᴴ = f(A, alg)` and returns the projector `Nᴴ' * Nᴴ` onto the nullspace, +which, unlike `Nᴴ` itself, is a differentiable function of `A`. +""" +lq_null_gauge_invariant_wrapper(f, A, alg) = (Nᴴ = f(A, alg); Nᴴ' * Nᴴ) + +""" + lq_null!_gauge_invariant_wrapper(f!, A, alg) + +In-place variant of [`lq_null_gauge_invariant_wrapper`](@ref), which zeros `A` after calling `f!`. +""" +lq_null!_gauge_invariant_wrapper(f!, A, alg) = (Nᴴ = call_and_zero!(f!, A, alg); Nᴴ' * Nᴴ) + function stabilize_eigvals!(D::AbstractVector) absD = collect(abs.(D)) p = invperm(sortperm(collect(absD))) # rank of abs(D) diff --git a/test/testsuite/enzyme/lq.jl b/test/testsuite/enzyme/lq.jl index e4aa8d8e3..9b56ce110 100644 --- a/test/testsuite/enzyme/lq.jl +++ b/test/testsuite/enzyme/lq.jl @@ -18,12 +18,14 @@ function test_enzyme_lq_compact( rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) - return @testset "lq_compact reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) + return @testset "lq_compact: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) A = instantiate_matrix(T, sz) alg = MatrixAlgebraKit.select_algorithm(lq_compact, A) LQ, ΔLQ = ad_lq_compact_setup(A) test_reverse(lq_compact, RT, (A, TA), (alg, Const); atol, rtol, output_tangent = ΔLQ, fdm) - test_reverse(call_and_zero!, RT, (lq_compact!, Const), (A, TA), (alg, Const); atol, rtol, output_tangent = ΔLQ, fdm) + test_reverse(call_and_zero!, RT, (lq_compact!, Const), (copy(A), TA), (alg, Const); atol, rtol, output_tangent = ΔLQ, fdm) + test_forward(lq_compact, RT, (A, TA), (alg, Const); atol, rtol, fdm) + test_forward(call_and_zero!, RT, (lq_compact!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm) end end @@ -32,7 +34,7 @@ function test_enzyme_lq_compact_rank_deficient( rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) - return @testset "lq_compact rank deficient A reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) + return @testset "lq_compact rank deficient A: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) A = instantiate_matrix(T, sz) m, n = size(A) r = min(m, n) - 5 @@ -40,7 +42,11 @@ function test_enzyme_lq_compact_rank_deficient( alg = MatrixAlgebraKit.select_algorithm(lq_compact, A) LQ, ΔLQ = ad_lq_compact_setup(A) test_reverse(lq_compact, RT, (A, TA), (alg, Const); atol, rtol, output_tangent = ΔLQ, fdm) - test_reverse(call_and_zero!, RT, (lq_compact!, Const), (A, TA), (alg, Const); atol, rtol, output_tangent = ΔLQ, fdm) + test_reverse(call_and_zero!, RT, (lq_compact!, Const), (copy(A), TA), (alg, Const); atol, rtol, output_tangent = ΔLQ, fdm) + # only the first r columns/rows of the isometric factor are differentiable + r = MatrixAlgebraKit.lq_rank(LQ[1]) + test_forward(lq_gauge_invariant_wrapper, RT, (lq_compact, Const), (A, TA), (alg, Const), (r, Const); atol, rtol, fdm) + test_forward(lq!_gauge_invariant_wrapper, RT, (lq_compact!, Const), (copy(A), TA), (alg, Const), (r, Const); atol, rtol, fdm) end end @@ -54,7 +60,11 @@ function test_enzyme_lq_full( alg = MatrixAlgebraKit.select_algorithm(lq_full, A) LQ, ΔLQ = ad_lq_full_setup(A) test_reverse(lq_full, RT, (A, TA), (alg, Const); atol, rtol, output_tangent = ΔLQ, fdm) - test_reverse(call_and_zero!, RT, (lq_full!, Const), (A, TA), (alg, Const); atol, rtol, output_tangent = ΔLQ, fdm) + test_reverse(call_and_zero!, RT, (lq_full!, Const), (copy(A), TA), (alg, Const); atol, rtol, output_tangent = ΔLQ, fdm) + # the extra columns/rows of the isometric factor are only determined up to a unitary rotation + r = min(size(A)...) + test_forward(lq_gauge_invariant_wrapper, RT, (lq_full, Const), (A, TA), (alg, Const), (r, Const); atol, rtol, fdm) + test_forward(lq!_gauge_invariant_wrapper, RT, (lq_full!, Const), (copy(A), TA), (alg, Const), (r, Const); atol, rtol, fdm) end end @@ -68,6 +78,9 @@ function test_enzyme_lq_null( alg = MatrixAlgebraKit.select_algorithm(lq_null, A) Nᴴ, ΔNᴴ = ad_lq_null_setup(A) test_reverse(lq_null, RT, (A, TA), (alg, Const); atol, rtol, output_tangent = ΔNᴴ) - test_reverse(call_and_zero!, RT, (lq_null!, Const), (A, TA), (alg, Const); atol, rtol, output_tangent = ΔNᴴ) + test_reverse(call_and_zero!, RT, (lq_null!, Const), (copy(A), TA), (alg, Const); atol, rtol, output_tangent = ΔNᴴ) + # the nullspace basis is only determined up to a unitary rotation + test_forward(lq_null_gauge_invariant_wrapper, RT, (lq_null, Const), (A, TA), (alg, Const); atol, rtol) + test_forward(lq_null!_gauge_invariant_wrapper, RT, (lq_null!, Const), (copy(A), TA), (alg, Const); atol, rtol) end end diff --git a/test/testsuite/enzyme/orthnull.jl b/test/testsuite/enzyme/orthnull.jl index 280e546eb..1119fe042 100644 --- a/test/testsuite/enzyme/orthnull.jl +++ b/test/testsuite/enzyme/orthnull.jl @@ -34,7 +34,9 @@ function test_enzyme_left_orth( alg = MatrixAlgebraKit.select_algorithm(left_orth!, A, :qr) VC, ΔVC = ad_left_orth_setup(A) test_reverse(left_orth, RT, (A, TA), (alg, Const); atol, rtol, fdm, output_tangent = ΔVC) - test_reverse(call_and_zero!, RT, (left_orth!, Const), (A, TA), (alg, Const); atol, rtol, fdm, output_tangent = ΔVC) + test_reverse(call_and_zero!, RT, (left_orth!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm, output_tangent = ΔVC) + test_forward(left_orth, RT, (A, TA), (alg, Const); atol, rtol, fdm) + test_forward(call_and_zero!, RT, (left_orth!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm) end if m >= n && !(T <: Diagonal) @@ -43,10 +45,10 @@ function test_enzyme_left_orth( alg = MatrixAlgebraKit.select_algorithm(left_orth!, A, :polar) VC, ΔVC = ad_left_orth_setup(A) test_reverse(left_orth, RT, (A, TA), (alg, Const); atol, rtol, fdm, output_tangent = ΔVC) - test_reverse(call_and_zero!, RT, (left_orth!, Const), (A, TA), (alg, Const); atol, rtol, fdm, output_tangent = ΔVC) + test_reverse(call_and_zero!, RT, (left_orth!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm, output_tangent = ΔVC) A = instantiate_matrix(T, sz) test_forward(left_orth, RT, (A, TA), (alg, Const); atol, rtol, fdm) - test_forward(call_and_zero!, RT, (left_orth!, Const), (A, TA), (alg, Const); atol, rtol, fdm) + test_forward(call_and_zero!, RT, (left_orth!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm) end end end @@ -71,7 +73,9 @@ function test_enzyme_right_orth( alg = MatrixAlgebraKit.select_algorithm(right_orth!, A, :lq) CVᴴ, ΔCVᴴ = ad_right_orth_setup(A) test_reverse(right_orth, RT, (A, TA), (alg, Const); atol, rtol, fdm, output_tangent = ΔCVᴴ) - test_reverse(call_and_zero!, RT, (right_orth!, Const), (A, TA), (alg, Const); atol, rtol, fdm, output_tangent = ΔCVᴴ) + test_reverse(call_and_zero!, RT, (right_orth!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm, output_tangent = ΔCVᴴ) + test_forward(right_orth, RT, (A, TA), (alg, Const); atol, rtol, fdm) + test_forward(call_and_zero!, RT, (right_orth!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm) end if m <= n && !(T <: Diagonal) @@ -80,10 +84,10 @@ function test_enzyme_right_orth( alg = MatrixAlgebraKit.select_algorithm(right_orth!, A, :polar) CVᴴ, ΔCVᴴ = ad_right_orth_setup(A) test_reverse(right_orth, RT, (A, TA), (alg, Const); atol, rtol, fdm, output_tangent = ΔCVᴴ) - test_reverse(call_and_zero!, RT, (right_orth!, Const), (A, TA), (alg, Const); atol, rtol, fdm, output_tangent = ΔCVᴴ) + test_reverse(call_and_zero!, RT, (right_orth!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm, output_tangent = ΔCVᴴ) A = instantiate_matrix(T, sz) test_forward(right_orth, RT, (A, TA), (alg, Const); atol, rtol, fdm) - test_forward(call_and_zero!, RT, (right_orth!, Const), (A, TA), (alg, Const); atol, rtol, fdm) + test_forward(call_and_zero!, RT, (right_orth!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm) end end end @@ -92,7 +96,7 @@ end """ test_enzyme_left_null(T, sz; rng, atol, rtol) -Test the Enzyme reverse-mode AD rule for `left_null` with the QR algorithm and its +Test the Enzyme forward- and reverse-mode AD rule for `left_null` with the QR algorithm and its in-place variant. """ function test_enzyme_left_null( @@ -100,13 +104,16 @@ function test_enzyme_left_null( rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) - return @testset "left_null reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) + return @testset "left_null: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) A = instantiate_matrix(T, sz) @testset "qr" begin alg = MatrixAlgebraKit.select_algorithm(left_null!, A, :qr) N, ΔN = ad_left_null_setup(A) test_reverse(left_null, RT, (A, TA), (alg, Const); output_tangent = ΔN, atol, rtol) - test_reverse(call_and_zero!, RT, (left_null!, Const), (A, TA), (alg, Const); output_tangent = ΔN, atol, rtol) + test_reverse(call_and_zero!, RT, (left_null!, Const), (copy(A), TA), (alg, Const); output_tangent = ΔN, atol, rtol) + # the nullspace basis is only determined up to a unitary rotation + test_forward(qr_null_gauge_invariant_wrapper, RT, (left_null, Const), (A, TA), (alg, Const); atol, rtol) + test_forward(qr_null!_gauge_invariant_wrapper, RT, (left_null!, Const), (copy(A), TA), (alg, Const); atol, rtol) end end end @@ -114,7 +121,7 @@ end """ test_enzyme_right_null(T, sz; rng, atol, rtol) -Test the Enzyme reverse-mode AD rule for `right_null` with the LQ algorithm and its +Test the Enzyme forward- and reverse-mode AD rule for `right_null` with the LQ algorithm and its in-place variant. """ function test_enzyme_right_null( @@ -122,13 +129,16 @@ function test_enzyme_right_null( rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T), fdm = enzyme_fdm(T) ) - return @testset "right_null reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) + return @testset "right_null: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,) A = instantiate_matrix(T, sz) @testset "lq" begin alg = MatrixAlgebraKit.select_algorithm(right_null!, A, :lq) Nᴴ, ΔNᴴ = ad_right_null_setup(A) test_reverse(right_null, RT, (A, TA), (alg, Const); output_tangent = ΔNᴴ, atol, rtol) - test_reverse(call_and_zero!, RT, (right_null!, Const), (A, TA), (alg, Const); output_tangent = ΔNᴴ, atol, rtol) + test_reverse(call_and_zero!, RT, (right_null!, Const), (copy(A), TA), (alg, Const); output_tangent = ΔNᴴ, atol, rtol) + # the nullspace basis is only determined up to a unitary rotation + test_forward(lq_null_gauge_invariant_wrapper, RT, (right_null, Const), (A, TA), (alg, Const); atol, rtol) + test_forward(lq_null!_gauge_invariant_wrapper, RT, (right_null!, Const), (copy(A), TA), (alg, Const); atol, rtol) end end end diff --git a/test/testsuite/enzyme/qr.jl b/test/testsuite/enzyme/qr.jl index 1d9f33a57..f3b8f0901 100644 --- a/test/testsuite/enzyme/qr.jl +++ b/test/testsuite/enzyme/qr.jl @@ -23,7 +23,9 @@ function test_enzyme_qr_compact( alg = MatrixAlgebraKit.select_algorithm(qr_compact, A) QR, ΔQR = ad_qr_compact_setup(A) test_reverse(qr_compact, RT, (A, TA), (alg, Const); atol, rtol, output_tangent = ΔQR, fdm) - test_reverse(call_and_zero!, RT, (qr_compact!, Const), (A, TA), (alg, Const); atol, rtol, output_tangent = ΔQR, fdm) + test_reverse(call_and_zero!, RT, (qr_compact!, Const), (copy(A), TA), (alg, Const); atol, rtol, output_tangent = ΔQR, fdm) + test_forward(qr_compact, RT, (A, TA), (alg, Const); atol, rtol, fdm) + test_forward(call_and_zero!, RT, (qr_compact!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm) end end @@ -40,7 +42,11 @@ function test_enzyme_qr_compact_rank_deficient( alg = MatrixAlgebraKit.select_algorithm(qr_compact, A) QR, ΔQR = ad_qr_compact_setup(A) test_reverse(qr_compact, RT, (A, TA), (alg, Const); atol, rtol, output_tangent = ΔQR, fdm) - test_reverse(call_and_zero!, RT, (qr_compact!, Const), (A, TA), (alg, Const); atol, rtol, output_tangent = ΔQR, fdm) + test_reverse(call_and_zero!, RT, (qr_compact!, Const), (copy(A), TA), (alg, Const); atol, rtol, output_tangent = ΔQR, fdm) + # only the first r columns/rows of the isometric factor are differentiable + r = MatrixAlgebraKit.qr_rank(QR[2]) + test_forward(qr_gauge_invariant_wrapper, RT, (qr_compact, Const), (A, TA), (alg, Const), (r, Const); atol, rtol, fdm) + test_forward(qr!_gauge_invariant_wrapper, RT, (qr_compact!, Const), (copy(A), TA), (alg, Const), (r, Const); atol, rtol, fdm) end end @@ -54,7 +60,11 @@ function test_enzyme_qr_full( alg = MatrixAlgebraKit.select_algorithm(qr_full, A) QR, ΔQR = ad_qr_full_setup(A) test_reverse(qr_full, RT, (A, TA), (alg, Const); atol, rtol, output_tangent = ΔQR, fdm) - test_reverse(call_and_zero!, RT, (qr_full!, Const), (A, TA), (alg, Const); atol, rtol, output_tangent = ΔQR, fdm) + test_reverse(call_and_zero!, RT, (qr_full!, Const), (copy(A), TA), (alg, Const); atol, rtol, output_tangent = ΔQR, fdm) + # the extra columns/rows of the isometric factor are only determined up to a unitary rotation + r = min(size(A)...) + test_forward(qr_gauge_invariant_wrapper, RT, (qr_full, Const), (A, TA), (alg, Const), (r, Const); atol, rtol, fdm) + test_forward(qr!_gauge_invariant_wrapper, RT, (qr_full!, Const), (copy(A), TA), (alg, Const), (r, Const); atol, rtol, fdm) end end @@ -68,6 +78,9 @@ function test_enzyme_qr_null( alg = MatrixAlgebraKit.select_algorithm(qr_null, A) N, ΔN = ad_qr_null_setup(A) test_reverse(qr_null, RT, (A, TA), (alg, Const); atol, rtol, output_tangent = ΔN) - test_reverse(call_and_zero!, RT, (qr_null!, Const), (A, TA), (alg, Const); atol, rtol, output_tangent = ΔN) + test_reverse(call_and_zero!, RT, (qr_null!, Const), (copy(A), TA), (alg, Const); atol, rtol, output_tangent = ΔN) + # the nullspace basis is only determined up to a unitary rotation + test_forward(qr_null_gauge_invariant_wrapper, RT, (qr_null, Const), (A, TA), (alg, Const); atol, rtol) + test_forward(qr_null!_gauge_invariant_wrapper, RT, (qr_null!, Const), (copy(A), TA), (alg, Const); atol, rtol) end end diff --git a/test/testsuite/mooncake/lq.jl b/test/testsuite/mooncake/lq.jl index 636f410b6..f947095fc 100644 --- a/test/testsuite/mooncake/lq.jl +++ b/test/testsuite/mooncake/lq.jl @@ -15,7 +15,7 @@ end """ test_mooncake_lq_compact(T, sz; rng, atol, rtol) -Test the Mooncake reverse-mode AD rule for `lq_compact` and its in-place variant. +Test the Mooncake forward- and reverse-mode AD rule for `lq_compact` and its in-place variant. """ function test_mooncake_lq_compact( T, sz; @@ -29,11 +29,11 @@ function test_mooncake_lq_compact( Mooncake.TestUtils.test_rule( rng, lq_compact, A, alg; - mode = Mooncake.ReverseMode, output_tangent, atol, rtol + output_tangent, atol, rtol ) Mooncake.TestUtils.test_rule( rng, call_and_zero!, lq_compact!, A, alg; - mode = Mooncake.ReverseMode, output_tangent, atol, rtol, is_primitive = false + output_tangent, atol, rtol, is_primitive = false ) A = instantiate_rank_deficient_matrix(T, sz) @@ -49,13 +49,25 @@ function test_mooncake_lq_compact( rng, call_and_zero!, lq_compact!, A, alg; mode = Mooncake.ReverseMode, output_tangent, atol, rtol, is_primitive = false ) + # only the first r rows of Q and columns of L are differentiable + if !(T <: Diagonal) # rank-deficient Diagonal does not have its first r rows independent + r = MatrixAlgebraKit.lq_rank(LQ[1]) + Mooncake.TestUtils.test_rule( + rng, lq_gauge_invariant_wrapper, lq_compact, A, alg, r; + mode = Mooncake.ForwardMode, atol, rtol, is_primitive = false + ) + Mooncake.TestUtils.test_rule( + rng, lq!_gauge_invariant_wrapper, lq_compact!, A, alg, r; + mode = Mooncake.ForwardMode, atol, rtol, is_primitive = false + ) + end end end """ test_mooncake_lq_full(T, sz; rng, atol, rtol) -Test the Mooncake reverse-mode AD rule for `lq_full` and its in-place variant. +Test the Mooncake forward- and reverse-mode AD rule for `lq_full` and its in-place variant. """ function test_mooncake_lq_full( T, sz; @@ -75,13 +87,23 @@ function test_mooncake_lq_full( rng, call_and_zero!, lq_full!, A, alg; mode = Mooncake.ReverseMode, output_tangent, atol, rtol, is_primitive = false ) + # the extra rows of Q are only determined up to a unitary rotation + r = min(size(A)...) + Mooncake.TestUtils.test_rule( + rng, lq_gauge_invariant_wrapper, lq_full, A, alg, r; + mode = Mooncake.ForwardMode, atol, rtol, is_primitive = false + ) + Mooncake.TestUtils.test_rule( + rng, lq!_gauge_invariant_wrapper, lq_full!, A, alg, r; + mode = Mooncake.ForwardMode, atol, rtol, is_primitive = false + ) end end """ test_mooncake_lq_null(T, sz; rng, atol, rtol) -Test the Mooncake reverse-mode AD rule for `lq_null` and its in-place variant. +Test the Mooncake forward- and reverse-mode AD rule for `lq_null` and its in-place variant. """ function test_mooncake_lq_null( T, sz; @@ -101,5 +123,14 @@ function test_mooncake_lq_null( rng, call_and_zero!, lq_null!, A, alg; mode = Mooncake.ReverseMode, output_tangent, atol, rtol, is_primitive = false ) + # the nullspace basis is only determined up to a unitary rotation + Mooncake.TestUtils.test_rule( + rng, lq_null_gauge_invariant_wrapper, lq_null, A, alg; + mode = Mooncake.ForwardMode, atol, rtol, is_primitive = false + ) + Mooncake.TestUtils.test_rule( + rng, lq_null!_gauge_invariant_wrapper, lq_null!, A, alg; + mode = Mooncake.ForwardMode, atol, rtol, is_primitive = false + ) end end diff --git a/test/testsuite/mooncake/orthnull.jl b/test/testsuite/mooncake/orthnull.jl index afbac9d99..f0285634a 100644 --- a/test/testsuite/mooncake/orthnull.jl +++ b/test/testsuite/mooncake/orthnull.jl @@ -35,11 +35,11 @@ function test_mooncake_left_orth( Mooncake.TestUtils.test_rule( rng, left_orth, A, alg; - mode = Mooncake.ReverseMode, output_tangent, is_primitive = false, atol, rtol + output_tangent, is_primitive = false, atol, rtol ) Mooncake.TestUtils.test_rule( rng, call_and_zero!, left_orth!, A, alg; - mode = Mooncake.ReverseMode, output_tangent, is_primitive = false, atol, rtol + output_tangent, is_primitive = false, atol, rtol ) end @@ -83,11 +83,11 @@ function test_mooncake_right_orth( Mooncake.TestUtils.test_rule( rng, right_orth, A, alg; - mode = Mooncake.ReverseMode, output_tangent, is_primitive = false, atol, rtol + output_tangent, is_primitive = false, atol, rtol ) Mooncake.TestUtils.test_rule( rng, call_and_zero!, right_orth!, A, alg; - mode = Mooncake.ReverseMode, output_tangent, is_primitive = false, atol, rtol + output_tangent, is_primitive = false, atol, rtol ) end @@ -113,7 +113,7 @@ end """ test_mooncake_left_null(T, sz; rng, atol, rtol) -Test the Mooncake reverse-mode AD rule for `left_null` with the QR algorithm and its +Test the Mooncake forward- and reverse-mode AD rule for `left_null` with the QR algorithm and its in-place variant. """ function test_mooncake_left_null( @@ -136,6 +136,15 @@ function test_mooncake_left_null( rng, call_and_zero!, left_null!, A, alg; mode = Mooncake.ReverseMode, output_tangent, is_primitive = false, atol, rtol ) + # the nullspace basis is only determined up to a unitary rotation + Mooncake.TestUtils.test_rule( + rng, qr_null_gauge_invariant_wrapper, left_null, A, alg; + mode = Mooncake.ForwardMode, is_primitive = false, atol, rtol + ) + Mooncake.TestUtils.test_rule( + rng, qr_null!_gauge_invariant_wrapper, left_null!, A, alg; + mode = Mooncake.ForwardMode, is_primitive = false, atol, rtol + ) end end end @@ -143,7 +152,7 @@ end """ test_mooncake_right_null(T, sz; rng, atol, rtol) -Test the Mooncake reverse-mode AD rule for `right_null` with the LQ algorithm and its +Test the Mooncake forward- and reverse-mode AD rule for `right_null` with the LQ algorithm and its in-place variant. """ function test_mooncake_right_null( @@ -166,6 +175,15 @@ function test_mooncake_right_null( rng, call_and_zero!, right_null!, A, alg; mode = Mooncake.ReverseMode, output_tangent, is_primitive = false, atol, rtol ) + # the nullspace basis is only determined up to a unitary rotation + Mooncake.TestUtils.test_rule( + rng, lq_null_gauge_invariant_wrapper, right_null, A, alg; + mode = Mooncake.ForwardMode, is_primitive = false, atol, rtol + ) + Mooncake.TestUtils.test_rule( + rng, lq_null!_gauge_invariant_wrapper, right_null!, A, alg; + mode = Mooncake.ForwardMode, is_primitive = false, atol, rtol + ) end end end diff --git a/test/testsuite/mooncake/qr.jl b/test/testsuite/mooncake/qr.jl index 2d917f3f3..830e3472b 100644 --- a/test/testsuite/mooncake/qr.jl +++ b/test/testsuite/mooncake/qr.jl @@ -15,7 +15,7 @@ end """ test_mooncake_qr_compact(T, sz; rng, atol, rtol) -Test the Mooncake reverse-mode AD rule for `qr_compact` and its in-place variant. +Test the Mooncake forward- and reverse-mode AD rule for `qr_compact` and its in-place variant. """ function test_mooncake_qr_compact( T, sz; @@ -29,11 +29,11 @@ function test_mooncake_qr_compact( Mooncake.TestUtils.test_rule( rng, qr_compact, A, alg; - mode = Mooncake.ReverseMode, output_tangent, atol, rtol + output_tangent, atol, rtol ) Mooncake.TestUtils.test_rule( rng, call_and_zero!, qr_compact!, A, alg; - mode = Mooncake.ReverseMode, output_tangent, atol, rtol, is_primitive = false + output_tangent, atol, rtol, is_primitive = false ) A = instantiate_rank_deficient_matrix(T, sz) @@ -49,13 +49,25 @@ function test_mooncake_qr_compact( rng, call_and_zero!, qr_compact!, A, alg; mode = Mooncake.ReverseMode, output_tangent, atol, rtol, is_primitive = false ) + # only the first r columns of Q and rows of R are differentiable + if !(T <: Diagonal) # rank-deficient Diagonal does not have its first r columns independent + r = MatrixAlgebraKit.qr_rank(QR[2]) + Mooncake.TestUtils.test_rule( + rng, qr_gauge_invariant_wrapper, qr_compact, A, alg, r; + mode = Mooncake.ForwardMode, atol, rtol, is_primitive = false + ) + Mooncake.TestUtils.test_rule( + rng, qr!_gauge_invariant_wrapper, qr_compact!, A, alg, r; + mode = Mooncake.ForwardMode, atol, rtol, is_primitive = false + ) + end end end """ test_mooncake_qr_full(T, sz; rng, atol, rtol) -Test the Mooncake reverse-mode AD rule for `qr_full` and its in-place variant. +Test the Mooncake forward- and reverse-mode AD rule for `qr_full` and its in-place variant. """ function test_mooncake_qr_full( T, sz; @@ -75,13 +87,23 @@ function test_mooncake_qr_full( rng, call_and_zero!, qr_full!, A, alg; mode = Mooncake.ReverseMode, output_tangent, atol, rtol, is_primitive = false ) + # the extra columns of Q are only determined up to a unitary rotation + r = min(size(A)...) + Mooncake.TestUtils.test_rule( + rng, qr_gauge_invariant_wrapper, qr_full, A, alg, r; + mode = Mooncake.ForwardMode, atol, rtol, is_primitive = false + ) + Mooncake.TestUtils.test_rule( + rng, qr!_gauge_invariant_wrapper, qr_full!, A, alg, r; + mode = Mooncake.ForwardMode, atol, rtol, is_primitive = false + ) end end """ test_mooncake_qr_null(T, sz; rng, atol, rtol) -Test the Mooncake reverse-mode AD rule for `qr_null` and its in-place variant. +Test the Mooncake forward- and reverse-mode AD rule for `qr_null` and its in-place variant. """ function test_mooncake_qr_null( T, sz; @@ -101,5 +123,14 @@ function test_mooncake_qr_null( rng, call_and_zero!, qr_null!, A, alg; mode = Mooncake.ReverseMode, output_tangent, atol, rtol, is_primitive = false ) + # the nullspace basis is only determined up to a unitary rotation + Mooncake.TestUtils.test_rule( + rng, qr_null_gauge_invariant_wrapper, qr_null, A, alg; + mode = Mooncake.ForwardMode, atol, rtol, is_primitive = false + ) + Mooncake.TestUtils.test_rule( + rng, qr_null!_gauge_invariant_wrapper, qr_null!, A, alg; + mode = Mooncake.ForwardMode, atol, rtol, is_primitive = false + ) end end From 5aa17906a311e42e298b77e0af96a48b22c43edc Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Fri, 25 Sep 2026 12:26:11 -0400 Subject: [PATCH 02/13] Calm Enzyme on 1.10 --- ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl b/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl index 90f93883b..45a121171 100644 --- a/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl +++ b/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl @@ -18,6 +18,7 @@ using Enzyme.EnzymeCore using Enzyme.EnzymeCore: EnzymeRules using LinearAlgebra +@inline EnzymeRules.inactive_type(::Type{Alg}) where {Alg <: MatrixAlgebraKit.Householder} = true @inline EnzymeRules.inactive_type(::Type{Alg}) where {Alg <: MatrixAlgebraKit.AbstractAlgorithm} = true @inline EnzymeRules.inactive_type(::Type{TS}) where {TS <: MatrixAlgebraKit.TruncationStrategy} = true @inline EnzymeRules.inactive(::typeof(MatrixAlgebraKit.select_algorithm), func::F, A::AbstractMatrix, alg::Alg) where {F, Alg} = true @@ -135,7 +136,7 @@ for (f, pf) in ( ::Type{RT}, A::Annotation, arg::Annotation, - alg::Const{<:MatrixAlgebraKit.AbstractAlgorithm}, + alg::Annotation{<:MatrixAlgebraKit.AbstractAlgorithm}, ) where {RT} A_is_arg1 = !isa(A, Const) && has_equal_storage(A.val, arg.val[1]) A_is_arg2 = !isa(A, Const) && has_equal_storage(A.val, arg.val[2]) @@ -215,7 +216,7 @@ for (f, pb, pf) in ( ::Type{RT}, A::Annotation, arg::Annotation, - alg::Const{<:MatrixAlgebraKit.AbstractAlgorithm}, + alg::Annotation{<:MatrixAlgebraKit.AbstractAlgorithm}, ) where {RT} # here, A IS directly used in the pushforward, and overwritten # in the primal call, so we MUST copy its value From e19beb5769aa08482a5cadae628d322f3d7d690a Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sat, 26 Sep 2026 10:45:14 -0400 Subject: [PATCH 03/13] Use Mooncake branch on CUDA for now --- test/Project.toml | 1 + 1 file changed, 1 insertion(+) diff --git a/test/Project.toml b/test/Project.toml index a518221ac..f60b2110f 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -22,6 +22,7 @@ Zygote = "e88e6eb3-aa80-5325-afca-941959d7151f" [sources] MatrixAlgebraKit = {path = ".."} +Mooncake = {url = "https://github.com/chalk-lab/Mooncake.jl.git", rev = "ksh/cu_range_getindex"} [compat] Aqua = "0.6, 0.7, 0.8" From 4281ef51229d7bbf6172f9fe6e6f77970b599b57 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sun, 27 Sep 2026 02:50:51 -0400 Subject: [PATCH 04/13] Fix lq --- src/pushforwards/lq.jl | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/src/pushforwards/lq.jl b/src/pushforwards/lq.jl index 80d1a59a5..ceb6acab2 100644 --- a/src/pushforwards/lq.jl +++ b/src/pushforwards/lq.jl @@ -28,7 +28,9 @@ function lq_pushforward!( (m, n) == size(ΔA) || throw(DimensionMismatch("size of ΔA ($(size(ΔA))) does not match size of L*Q ($m, $n)")) Q₁ = view(Q, 1:p, :) - L₁₁ = LowerTriangular(view(L, 1:p, 1:p)) + # Julia 1.13 `ldiv!` checks `istriu` on the parent, which + # falls back to scalar indexing for a view of a GPU array + L₁₁ = LowerTriangular(L[1:p, 1:p]) L₂₁ = view(L, (p + 1):m, 1:p) ΔA₁ = view(ΔA, 1:p, :) From 30e7344accc8d53852556361ce0522c1eb6c1c30 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sun, 27 Sep 2026 09:29:24 +0200 Subject: [PATCH 05/13] Apply batched suggestions from code review Co-authored-by: Jutho --- src/pushforwards/lq.jl | 24 ++++++++++++------------ src/pushforwards/qr.jl | 6 +++--- 2 files changed, 15 insertions(+), 15 deletions(-) diff --git a/src/pushforwards/lq.jl b/src/pushforwards/lq.jl index ceb6acab2..07d20ea41 100644 --- a/src/pushforwards/lq.jl +++ b/src/pushforwards/lq.jl @@ -36,21 +36,21 @@ function lq_pushforward!( ΔA₁ = view(ΔA, 1:p, :) ΔA₂ = view(ΔA, (p + 1):m, :) - ΔQ₁ = L₁₁ \ ΔA₁ - ΔQ₁Q₁ᴴ = ΔQ₁ * Q₁' - M = ΔQ₁Q₁ᴴ + ΔQ₁Q₁ᴴ' - diagview(M) ./= 2 - view(M, uppertriangularind(M)) .= zero(eltype(M)) - ΔL₁₁ = L₁₁ * M - ΔQ₁ = mul!(ΔQ₁, M, Q₁, -1, 1) - ΔL₂₁ = ΔA₂ * Q₁' - ΔL₂₁ = mul!(ΔL₂₁, L₂₁, Q₁ * ΔQ₁', 1, 1) - zero!(ΔL) zero!(ΔQ) view(ΔQ, 1:p, :) .= ΔQ₁ view(ΔL, 1:p, 1:p) .= ΔL₁₁ view(ΔL, (p + 1):m, 1:p) .= ΔL₂₁ + + ΔQ₁ = ldiv!(L₁₁, copy!(ΔQ₁, ΔA₁)) + ΔQ₁Q₁ᴴ = ΔQ₁ * Q₁' + M = ΔQ₁Q₁ᴴ + ΔQ₁Q₁ᴴ' + diagview(M) ./= 2 + view(M, uppertriangularind(M)) .= zero(eltype(M)) + ΔL₁₁ = mul!!(ΔL₁₁, L₁₁, M) + ΔQ₁ = mul!(ΔQ₁, M, Q₁, -1, 1) + ΔL₂₁ = mul!(ΔL₂₁, ΔA₂, Q₁') + ΔL₂₁ = mul!(ΔL₂₁, L₂₁, ΔQ₁ * Q₁', -1, 1) if p == minmn && size(Q, 1) > minmn Q₃ = view(Q, (minmn + 1):size(Q, 1), :) ΔQ₃ = view(ΔQ, (minmn + 1):size(Q, 1), :) @@ -73,6 +73,6 @@ function lq_null_pushforward!(ΔA, A, Nᴴ, ΔNᴴ; kwargs...) return ΔNᴴ end L, Q = lq_compact(A; positive = true) - X = ldiv!(LowerTriangular(L), ΔA * Nᴴ') - return mul!(ΔNᴴ, X', Q, -1, 0) + ΔQNᴴ = ldiv!(LowerTriangular(L), ΔA * Nᴴ') + return mul!(ΔNᴴ, ΔQNᴴ', Q, -1, 0) end diff --git a/src/pushforwards/qr.jl b/src/pushforwards/qr.jl index 67631d9c6..6190f3354 100644 --- a/src/pushforwards/qr.jl +++ b/src/pushforwards/qr.jl @@ -40,7 +40,7 @@ function qr_pushforward!( ΔR₁₁ = M * R₁₁ ΔQ₁ = mul!(ΔQ₁, Q₁, M, -1, 1) ΔR₁₂ = Q₁' * ΔA₂ - ΔR₁₂ = mul!(ΔR₁₂, ΔQ₁' * Q₁, R₁₂, 1, 1) + ΔR₁₂ = mul!(ΔR₁₂, Q₁' * ΔQ₁, R₁₂, -1, 1) zero!(ΔQ) zero!(ΔR) @@ -70,6 +70,6 @@ function qr_null_pushforward!(ΔA, A, N, ΔN; kwargs...) return ΔN end Q, R = qr_compact(A; positive = true) - X = ldiv!(UpperTriangular(R)', ΔA' * N) - return mul!(ΔN, Q, X, -1, 0) + NᴴΔQ = rdiv!(N' * ΔA, UpperTriangular(R)) + return mul!(ΔN, Q, NᴴΔQ', -1, 0) end From 76221d693c1d9947efe777d7e49c5e36f1357bf8 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sun, 27 Sep 2026 09:31:00 +0200 Subject: [PATCH 06/13] Formatter --- src/pushforwards/lq.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/pushforwards/lq.jl b/src/pushforwards/lq.jl index 07d20ea41..207baade7 100644 --- a/src/pushforwards/lq.jl +++ b/src/pushforwards/lq.jl @@ -41,7 +41,7 @@ function lq_pushforward!( view(ΔQ, 1:p, :) .= ΔQ₁ view(ΔL, 1:p, 1:p) .= ΔL₁₁ view(ΔL, (p + 1):m, 1:p) .= ΔL₂₁ - + ΔQ₁ = ldiv!(L₁₁, copy!(ΔQ₁, ΔA₁)) ΔQ₁Q₁ᴴ = ΔQ₁ * Q₁' M = ΔQ₁Q₁ᴴ + ΔQ₁Q₁ᴴ' From a6135058fe854ddc37e6cf4d46f9ee016aa09f8b Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sun, 27 Sep 2026 10:15:04 +0200 Subject: [PATCH 07/13] Fix typos from suggestions --- src/pushforwards/lq.jl | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/src/pushforwards/lq.jl b/src/pushforwards/lq.jl index 207baade7..39bedbe20 100644 --- a/src/pushforwards/lq.jl +++ b/src/pushforwards/lq.jl @@ -28,10 +28,13 @@ function lq_pushforward!( (m, n) == size(ΔA) || throw(DimensionMismatch("size of ΔA ($(size(ΔA))) does not match size of L*Q ($m, $n)")) Q₁ = view(Q, 1:p, :) + ΔQ₁ = view(ΔQ, 1:p, :) # Julia 1.13 `ldiv!` checks `istriu` on the parent, which # falls back to scalar indexing for a view of a GPU array L₁₁ = LowerTriangular(L[1:p, 1:p]) + ΔL₁₁ = LowerTriangular(ΔL[1:p, 1:p]) L₂₁ = view(L, (p + 1):m, 1:p) + ΔL₂₁ = view(ΔL, (p + 1):m, 1:p) ΔA₁ = view(ΔA, 1:p, :) ΔA₂ = view(ΔA, (p + 1):m, :) @@ -47,7 +50,7 @@ function lq_pushforward!( M = ΔQ₁Q₁ᴴ + ΔQ₁Q₁ᴴ' diagview(M) ./= 2 view(M, uppertriangularind(M)) .= zero(eltype(M)) - ΔL₁₁ = mul!!(ΔL₁₁, L₁₁, M) + ΔL₁₁ = mul!(ΔL₁₁, L₁₁, M) ΔQ₁ = mul!(ΔQ₁, M, Q₁, -1, 1) ΔL₂₁ = mul!(ΔL₂₁, ΔA₂, Q₁') ΔL₂₁ = mul!(ΔL₂₁, L₂₁, ΔQ₁ * Q₁', -1, 1) From 847b5c1e732f2e00db996ffc1983293213d4186e Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sun, 27 Sep 2026 17:04:17 -0400 Subject: [PATCH 08/13] Fix in case of aliases --- src/pushforwards/lq.jl | 23 +++++++++++------------ 1 file changed, 11 insertions(+), 12 deletions(-) diff --git a/src/pushforwards/lq.jl b/src/pushforwards/lq.jl index 39bedbe20..3d327de9e 100644 --- a/src/pushforwards/lq.jl +++ b/src/pushforwards/lq.jl @@ -28,32 +28,31 @@ function lq_pushforward!( (m, n) == size(ΔA) || throw(DimensionMismatch("size of ΔA ($(size(ΔA))) does not match size of L*Q ($m, $n)")) Q₁ = view(Q, 1:p, :) - ΔQ₁ = view(ΔQ, 1:p, :) # Julia 1.13 `ldiv!` checks `istriu` on the parent, which # falls back to scalar indexing for a view of a GPU array L₁₁ = LowerTriangular(L[1:p, 1:p]) - ΔL₁₁ = LowerTriangular(ΔL[1:p, 1:p]) L₂₁ = view(L, (p + 1):m, 1:p) - ΔL₂₁ = view(ΔL, (p + 1):m, 1:p) ΔA₁ = view(ΔA, 1:p, :) ΔA₂ = view(ΔA, (p + 1):m, :) - zero!(ΔL) - zero!(ΔQ) - view(ΔQ, 1:p, :) .= ΔQ₁ - view(ΔL, 1:p, 1:p) .= ΔL₁₁ - view(ΔL, (p + 1):m, 1:p) .= ΔL₂₁ - - ΔQ₁ = ldiv!(L₁₁, copy!(ΔQ₁, ΔA₁)) + # compute everything from ΔA before writing to ΔL and ΔQ, which may alias it + # (e.g. `lq_compact!` of a `Diagonal` returns `Q === A`) + ΔQ₁ = L₁₁ \ ΔA₁ ΔQ₁Q₁ᴴ = ΔQ₁ * Q₁' M = ΔQ₁Q₁ᴴ + ΔQ₁Q₁ᴴ' diagview(M) ./= 2 view(M, uppertriangularind(M)) .= zero(eltype(M)) - ΔL₁₁ = mul!(ΔL₁₁, L₁₁, M) + ΔL₁₁ = L₁₁ * M ΔQ₁ = mul!(ΔQ₁, M, Q₁, -1, 1) - ΔL₂₁ = mul!(ΔL₂₁, ΔA₂, Q₁') + ΔL₂₁ = ΔA₂ * Q₁' ΔL₂₁ = mul!(ΔL₂₁, L₂₁, ΔQ₁ * Q₁', -1, 1) + + zero!(ΔL) + zero!(ΔQ) + view(ΔQ, 1:p, :) .= ΔQ₁ + view(ΔL, 1:p, 1:p) .= ΔL₁₁ + view(ΔL, (p + 1):m, 1:p) .= ΔL₂₁ if p == minmn && size(Q, 1) > minmn Q₃ = view(Q, (minmn + 1):size(Q, 1), :) ΔQ₃ = view(ΔQ, (minmn + 1):size(Q, 1), :) From a5a0eeb9bcb564db48c21a3f0be7d148bc564073 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 30 Sep 2026 07:15:42 +0200 Subject: [PATCH 09/13] Update Project.tomls --- Project.toml | 2 +- test/Project.toml | 1 - 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/Project.toml b/Project.toml index afbc3531d..d8c262b39 100644 --- a/Project.toml +++ b/Project.toml @@ -37,5 +37,5 @@ GenericLinearAlgebra = "0.3.19, 0.4" GenericSchur = "0.5.6" LinearAlgebra = "1" PrecompileTools = "1" -Mooncake = "0.5.27" +Mooncake = "0.5.61" julia = "1.10" diff --git a/test/Project.toml b/test/Project.toml index f60b2110f..a518221ac 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -22,7 +22,6 @@ Zygote = "e88e6eb3-aa80-5325-afca-941959d7151f" [sources] MatrixAlgebraKit = {path = ".."} -Mooncake = {url = "https://github.com/chalk-lab/Mooncake.jl.git", rev = "ksh/cu_range_getindex"} [compat] Aqua = "0.6, 0.7, 0.8" From d723c58f2033be0a241c7118b214bcc809298566 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 30 Sep 2026 12:48:27 +0200 Subject: [PATCH 10/13] Remove extraneous inactive_type line --- ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl | 1 - 1 file changed, 1 deletion(-) diff --git a/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl b/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl index 45a121171..944833237 100644 --- a/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl +++ b/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl @@ -18,7 +18,6 @@ using Enzyme.EnzymeCore using Enzyme.EnzymeCore: EnzymeRules using LinearAlgebra -@inline EnzymeRules.inactive_type(::Type{Alg}) where {Alg <: MatrixAlgebraKit.Householder} = true @inline EnzymeRules.inactive_type(::Type{Alg}) where {Alg <: MatrixAlgebraKit.AbstractAlgorithm} = true @inline EnzymeRules.inactive_type(::Type{TS}) where {TS <: MatrixAlgebraKit.TruncationStrategy} = true @inline EnzymeRules.inactive(::typeof(MatrixAlgebraKit.select_algorithm), func::F, A::AbstractMatrix, alg::Alg) where {F, Alg} = true From 1b5b4a80982359b1d24a859647b85eb4d2ed5849 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 30 Sep 2026 14:12:34 +0200 Subject: [PATCH 11/13] Remove unnecessary reinstantiation of A --- test/testsuite/enzyme/orthnull.jl | 2 -- 1 file changed, 2 deletions(-) diff --git a/test/testsuite/enzyme/orthnull.jl b/test/testsuite/enzyme/orthnull.jl index 1119fe042..05a1640ea 100644 --- a/test/testsuite/enzyme/orthnull.jl +++ b/test/testsuite/enzyme/orthnull.jl @@ -46,7 +46,6 @@ function test_enzyme_left_orth( VC, ΔVC = ad_left_orth_setup(A) test_reverse(left_orth, RT, (A, TA), (alg, Const); atol, rtol, fdm, output_tangent = ΔVC) test_reverse(call_and_zero!, RT, (left_orth!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm, output_tangent = ΔVC) - A = instantiate_matrix(T, sz) test_forward(left_orth, RT, (A, TA), (alg, Const); atol, rtol, fdm) test_forward(call_and_zero!, RT, (left_orth!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm) end @@ -85,7 +84,6 @@ function test_enzyme_right_orth( CVᴴ, ΔCVᴴ = ad_right_orth_setup(A) test_reverse(right_orth, RT, (A, TA), (alg, Const); atol, rtol, fdm, output_tangent = ΔCVᴴ) test_reverse(call_and_zero!, RT, (right_orth!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm, output_tangent = ΔCVᴴ) - A = instantiate_matrix(T, sz) test_forward(right_orth, RT, (A, TA), (alg, Const); atol, rtol, fdm) test_forward(call_and_zero!, RT, (right_orth!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm) end From 44a2308dda52872a032ba91b49cdcd821cbbcf76 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 30 Sep 2026 14:15:51 +0200 Subject: [PATCH 12/13] Add todo comment to lq --- src/pushforwards/lq.jl | 2 ++ 1 file changed, 2 insertions(+) diff --git a/src/pushforwards/lq.jl b/src/pushforwards/lq.jl index 3d327de9e..9d1f82eed 100644 --- a/src/pushforwards/lq.jl +++ b/src/pushforwards/lq.jl @@ -38,6 +38,8 @@ function lq_pushforward!( # compute everything from ΔA before writing to ΔL and ΔQ, which may alias it # (e.g. `lq_compact!` of a `Diagonal` returns `Q === A`) + # TODO rework this into two versions, one for `Q === A` and one for + # `Q !== A`, to minimize allocations ΔQ₁ = L₁₁ \ ΔA₁ ΔQ₁Q₁ᴴ = ΔQ₁ * Q₁' M = ΔQ₁Q₁ᴴ + ΔQ₁Q₁ᴴ' From 3f9e51bfe83d617b94525534dee50454756b413b Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 30 Sep 2026 14:16:41 +0200 Subject: [PATCH 13/13] Add a TODO comment to qr pushforward also --- src/pushforwards/qr.jl | 2 ++ 1 file changed, 2 insertions(+) diff --git a/src/pushforwards/qr.jl b/src/pushforwards/qr.jl index 6190f3354..a4fd7fc16 100644 --- a/src/pushforwards/qr.jl +++ b/src/pushforwards/qr.jl @@ -25,6 +25,8 @@ function qr_pushforward!( p = qr_rank(R; rank_atol) (m, n) == size(ΔA) || throw(DimensionMismatch("size of ΔA ($(size(ΔA))) does not match size of Q*R ($m, $n)")) + # TODO rework this into two versions, one for `Q === A` and one for + # `Q !== A`, to minimize allocations Q₁ = view(Q, :, 1:p) R₁₁ = UpperTriangular(view(R, 1:p, 1:p)) R₁₂ = view(R, 1:p, (p + 1):n)