Skip to content
2 changes: 1 addition & 1 deletion Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
40 changes: 36 additions & 4 deletions ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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!
Expand Down Expand Up @@ -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(
Expand All @@ -129,7 +135,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])
Expand All @@ -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(
Expand Down Expand Up @@ -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::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
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

Expand Down
34 changes: 34 additions & 0 deletions ext/MatrixAlgebraKitMooncakeExt/MatrixAlgebraKitMooncakeExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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!
Expand Down Expand Up @@ -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!),
)
Expand Down Expand Up @@ -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),
Expand Down
2 changes: 2 additions & 0 deletions src/MatrixAlgebraKit.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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")

Expand Down
82 changes: 82 additions & 0 deletions src/pushforwards/lq.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,82 @@
"""
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, :)
# 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, :)
ΔA₂ = view(ΔA, (p + 1):m, :)

# 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₁ᴴ'
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)
ΔQNᴴ = ldiv!(LowerTriangular(L), ΔA * Nᴴ')
return mul!(ΔNᴴ, ΔQNᴴ', Q, -1, 0)
end
77 changes: 77 additions & 0 deletions src/pushforwards/qr.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,77 @@
"""
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)"))

# 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)

Δ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₁₂
Comment thread
kshyatt marked this conversation as resolved.
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)
NᴴΔQ = rdiv!(N' * ΔA, UpperTriangular(R))
return mul!(ΔN, Q, NᴴΔQ', -1, 0)
end
78 changes: 78 additions & 0 deletions test/testsuite/ad_utils.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading
Loading