Skip to content

Commit f22789c

Browse files
committed
simplify implementations
1 parent 8ba9c5a commit f22789c

2 files changed

Lines changed: 20 additions & 82 deletions

File tree

ext/MatrixAlgebraKitChainRulesCoreExt.jl

Lines changed: 6 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -274,46 +274,22 @@ function ChainRulesCore.rrule(::typeof(right_polar!), A, PWᴴ, alg)
274274
return PWᴴ, right_polar_pullback
275275
end
276276

277-
function ChainRulesCore.rrule(::typeof(project_hermitian!), A, Aₕ, alg)
278-
Ac = copy_input(project_hermitian, A)
279-
Aₕ = project_hermitian!(Ac, Aₕ, alg)
277+
function ChainRulesCore.rrule(::typeof(project_hermitian), A, alg)
278+
Aₕ = project_hermitian(A, alg)
280279
function project_hermitian_pullback(ΔAₕ)
281280
ΔA = project_hermitian(unthunk(ΔAₕ))
282-
return NoTangent(), ΔA, ZeroTangent(), NoTangent()
283-
end
284-
function project_hermitian_pullback(::ZeroTangent)
285-
return NoTangent(), ZeroTangent(), ZeroTangent(), NoTangent()
281+
return NoTangent(), ΔA, NoTangent()
286282
end
287283
return Aₕ, project_hermitian_pullback
288284
end
289285

290-
function ChainRulesCore.rrule(::typeof(project_antihermitian!), A, Aₐ, alg)
291-
Ac = copy_input(project_antihermitian, A)
292-
Aₐ = project_antihermitian!(Ac, Aₐ, alg)
286+
function ChainRulesCore.rrule(::typeof(project_antihermitian), A, alg)
287+
Aₐ = project_antihermitian(A, alg)
293288
function project_antihermitian_pullback(ΔAₐ)
294289
ΔA = project_antihermitian(unthunk(ΔAₐ))
295-
return NoTangent(), ΔA, ZeroTangent(), NoTangent()
296-
end
297-
function project_antihermitian_pullback(::ZeroTangent)
298-
return NoTangent(), ZeroTangent(), ZeroTangent(), NoTangent()
290+
return NoTangent(), ΔA, NoTangent()
299291
end
300292
return Aₐ, project_antihermitian_pullback
301293
end
302294

303-
function ChainRulesCore.rrule(::typeof(project_isometric!), A, W, alg)
304-
Ac = copy_input(project_isometric, A)
305-
# Compute the full polar decomposition to cache P for the pullback
306-
WP = left_polar!(Ac, (similar(W), similar(W, size(W, 2), size(W, 2))), alg)
307-
W_out = copy!(W, WP[1])
308-
function project_isometric_pullback(ΔW)
309-
ΔA = zero(A)
310-
MatrixAlgebraKit.left_polar_pullback!(ΔA, A, WP, (unthunk(ΔW), nothing))
311-
return NoTangent(), ΔA, ZeroTangent(), NoTangent()
312-
end
313-
function project_isometric_pullback(::ZeroTangent)
314-
return NoTangent(), ZeroTangent(), ZeroTangent(), NoTangent()
315-
end
316-
return W_out, project_isometric_pullback
317-
end
318-
319295
end

ext/MatrixAlgebraKitMooncakeExt/MatrixAlgebraKitMooncakeExt.jl

Lines changed: 14 additions & 52 deletions
Original file line numberDiff line numberDiff line change
@@ -778,82 +778,44 @@ function Mooncake.rrule!!(::CoDual{typeof(svd_trunc_no_error)}, A_dA::CoDual, al
778778
return USVᴴtrunc_dUSVᴴtrunc, svd_trunc_adjoint
779779
end
780780

781-
# single-output projections: project_hermitian!, project_antihermitian!
782781
# single-output projections: project_hermitian!, project_antihermitian!
783782
for (f!, f, adj) in (
784783
(:project_hermitian!, :project_hermitian, :project_hermitian_adjoint),
785784
(:project_antihermitian!, :project_antihermitian, :project_antihermitian_adjoint),
786785
)
787786
@eval begin
788-
@is_primitive Mooncake.DefaultCtx Mooncake.ReverseMode Tuple{typeof($f!), Any, Any, MatrixAlgebraKit.AbstractAlgorithm}
787+
@is_primitive DefaultCtx Mooncake.ReverseMode Tuple{typeof($f!), Any, Any, MatrixAlgebraKit.AbstractAlgorithm}
789788
function Mooncake.rrule!!(f_df::CoDual{typeof($f!)}, A_dA::CoDual, arg_darg::CoDual, alg_dalg::CoDual{<:MatrixAlgebraKit.AbstractAlgorithm})
790789
A, dA = arrayify(A_dA)
791-
Ac = copy(A)
792-
arg, darg = arrayify(arg_darg)
790+
arg, darg = A_dA === arg_darg ? (A, dA) : arrayify(arg_darg)
793791
argc = copy(arg)
794-
$f!(A, arg, Mooncake.primal(alg_dalg))
792+
arg = $f!(A, arg, Mooncake.primal(alg_dalg))
793+
795794
function $adj(::NoRData)
796-
copy!(A, Ac)
797795
dA .+= $f(darg)
796+
dA === darg || zero!(darg)
798797
copy!(arg, argc)
799-
zero!(darg)
800798
return NoRData(), NoRData(), NoRData(), NoRData()
801799
end
802800
return arg_darg, $adj
803801
end
804-
@is_primitive Mooncake.DefaultCtx Mooncake.ReverseMode Tuple{typeof($f), Any, MatrixAlgebraKit.AbstractAlgorithm}
802+
803+
@is_primitive DefaultCtx Mooncake.ReverseMode Tuple{typeof($f), Any, MatrixAlgebraKit.AbstractAlgorithm}
805804
function Mooncake.rrule!!(f_df::CoDual{typeof($f)}, A_dA::CoDual, alg_dalg::CoDual{<:MatrixAlgebraKit.AbstractAlgorithm})
806805
A, dA = arrayify(A_dA)
807806
output = $f(A, Mooncake.primal(alg_dalg))
808-
output_codual = CoDual(output, Mooncake.zero_tangent(output))
807+
output_doutput = Mooncake.zero_fcodual(output)
808+
809+
doutput = last(arrayify(output_doutput))
809810
function $adj(::NoRData)
810-
arg, darg = arrayify(output_codual)
811-
dA .+= $f(darg)
812-
zero!(darg)
813-
return NoRData(), NoRData(), NoRData()
811+
# TODO: need accumulating projection to avoid intermediate here
812+
dA .+= $f(doutput)
813+
return ntuple(Returns(NoRData(), 3))
814814
end
815+
815816
return output_codual, $adj
816817
end
817818
end
818819
end
819820

820-
# project_isometric! needs special handling: compute full polar decomposition
821-
@is_primitive Mooncake.DefaultCtx Mooncake.ReverseMode Tuple{typeof(project_isometric!), Any, Any, MatrixAlgebraKit.AbstractAlgorithm}
822-
function Mooncake.rrule!!(f_df::CoDual{typeof(project_isometric!)}, A_dA::CoDual, W_dW::CoDual, alg_dalg::CoDual{<:MatrixAlgebraKit.AbstractAlgorithm})
823-
A, dA = arrayify(A_dA)
824-
W, dW = arrayify(W_dW)
825-
Ac = copy(A)
826-
Wc = copy(W)
827-
# Compute the full polar decomposition for the pullback
828-
m, n = size(A)
829-
P = similar(A, n, n)
830-
WP = left_polar!(copy(A), (copy(W), P), Mooncake.primal(alg_dalg))
831-
copy!(W, WP[1])
832-
function project_isometric_adjoint(::NoRData)
833-
copy!(A, Ac)
834-
left_polar_pullback!(dA, A, WP, (dW, nothing))
835-
copy!(W, Wc)
836-
zero!(dW)
837-
return NoRData(), NoRData(), NoRData(), NoRData()
838-
end
839-
return W_dW, project_isometric_adjoint
840-
end
841-
842-
@is_primitive Mooncake.DefaultCtx Mooncake.ReverseMode Tuple{typeof(project_isometric), Any, MatrixAlgebraKit.AbstractAlgorithm}
843-
function Mooncake.rrule!!(f_df::CoDual{typeof(project_isometric)}, A_dA::CoDual, alg_dalg::CoDual{<:MatrixAlgebraKit.AbstractAlgorithm})
844-
A, dA = arrayify(A_dA)
845-
alg = Mooncake.primal(alg_dalg)
846-
# Compute the full polar decomposition for the pullback
847-
WP = left_polar(A, alg)
848-
W_out = WP[1]
849-
output_codual = CoDual(W_out, Mooncake.zero_tangent(W_out))
850-
function project_isometric_adjoint(::NoRData)
851-
W, dW = arrayify(output_codual)
852-
left_polar_pullback!(dA, A, WP, (dW, nothing))
853-
zero!(dW)
854-
return NoRData(), NoRData(), NoRData()
855-
end
856-
return output_codual, project_isometric_adjoint
857-
end
858-
859821
end

0 commit comments

Comments
 (0)