@@ -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
779779end
780780
781- # single-output projections: project_hermitian!, project_antihermitian!
782781# single-output projections: project_hermitian!, project_antihermitian!
783782for (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
818819end
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-
859821end
0 commit comments