diff --git a/src/common/pullbacks.jl b/src/common/pullbacks.jl index 053db5335..43bba4c45 100644 --- a/src/common/pullbacks.jl +++ b/src/common/pullbacks.jl @@ -35,3 +35,20 @@ iterating over `ind`, so that this also works for an `ind` that lives on a devic """ is_leading_index(ind::AbstractRange, p::Int) = ind == 1:p is_leading_index(ind::AbstractVector, p::Int) = length(ind) == p && all(ind .== 1:p) + +""" + antihermitian_columns!(X, ind) + +Given the columns `X = M[:, ind]` of a square matrix `M` that is nonzero only in its columns +`ind`, overwrite `X` with the same columns of the antihermitian part `(M - M') / 2` and return it. +Here `M'` is nonzero only in the rows `ind`, so it only contributes to the square block +`X[ind, :] = M[ind, ind]` on the diagonal of `M`. +""" +function antihermitian_columns!(X, ind) + # NOTE: all columns in order (e.g. from `ind = Colon()` in the pullback): the original in-place projection + is_leading_index(ind, size(X, 1)) && return project_antihermitian!(X) + Xdiag = project_antihermitian!(X[ind, :]) # the diagonal block M[ind, ind] + X ./= 2 + X[ind, :] .= Xdiag + return X +end diff --git a/src/pullbacks/eigh.jl b/src/pullbacks/eigh.jl index 713c04b5c..bd34c6b60 100755 --- a/src/pullbacks/eigh.jl +++ b/src/pullbacks/eigh.jl @@ -4,45 +4,39 @@ function check_and_prepare_eigh_cotangents( gauge_atol::Real = default_pullback_gauge_atol(ΔDmat, ΔV) ) + # only the columns ind of VᴴΔV and VᴴAΔV are computed; their rows ind follow by antihermiticity n, p = size(V) - indD = select_indices(axes(D, 1), ind) - indV = select_indices(axes(V, 2), ind) + ind′ = select_indices(axes(D, 1), ind) + k = length(ind′) if !iszerotangent(ΔV) n == size(ΔV, 1) || throw(DimensionMismatch()) - length(indV) == size(ΔV, 2) || throw(DimensionMismatch()) - if is_leading_index(indV, p) - ΔV₁ = copy(ΔV) - else - ΔV₁ = zero(V) - ΔV₁[:, indV] = ΔV - end - VᴴΔV₁ = V' * ΔV₁ + k == size(ΔV, 2) || throw(DimensionMismatch()) + VᴴΔV₁ = V' * ΔV if p == n - ΔV₊ = zero!(ΔV₁) + ΔV₊ = zero(ΔV) else - ΔV₊ = mul!(ΔV₁, V, VᴴΔV₁, -1, 1) + ΔV₊ = mul!(copy(ΔV), V, VᴴΔV₁, -1, 1) end - aVᴴΔV₁ = project_antihermitian!(VᴴΔV₁) + aVᴴΔV₁ = antihermitian_columns!(VᴴΔV₁, ind′) else ΔV₊ = nothing - aVᴴΔV₁ = zero!(similar(V, (p, p))) + aVᴴΔV₁ = zero!(similar(V, (p, k))) end - bc = Base.broadcasted(transpose(D), D, aVᴴΔV₁) do d₁, d₂, v - return abs(d₁ - d₂) < degeneracy_atol ? v : zero(v) - end - Δgauge = maximum(abs, Base.Broadcast.instantiate(bc); init = abs(zero(eltype(D)))) + Dₖ = view(D, ind′) + gauge_part = (abs.(transpose(Dₖ) .- D) .< degeneracy_atol) .* aVᴴΔV₁ + Δgauge = maximum(abs, gauge_part; init = abs(zero(eltype(D)))) Δgauge ≤ gauge_atol || @warn "`eigh` cotangents sensitive to gauge choice: (|Δgauge| = $Δgauge)" - aVᴴΔV₁ .*= inv_safe.(D' .- D, degeneracy_atol) + aVᴴΔV₁ .*= inv_safe.(transpose(Dₖ) .- D, degeneracy_atol) VᴴAΔV = aVᴴΔV₁ if !iszerotangent(ΔDmat) ΔD = diagview(ΔDmat) - length(indD) == length(ΔD) || throw(DimensionMismatch()) - VᴴAΔV[select_indices(diagind(VᴴAΔV), indD)] .+= real.(ΔD) + k == length(ΔD) || throw(DimensionMismatch()) + diagview(view(VᴴAΔV, ind′, :)) .+= real.(ΔD) else ΔD = nothing end @@ -86,11 +80,29 @@ function eigh_pullback!( iszero(n) && return ΔA ΔDmat, ΔV = ΔDV - VᴴΔAV, = check_and_prepare_eigh_cotangents( + VᴴΔAVₖ, = check_and_prepare_eigh_cotangents( D, V, ΔDmat, ΔV, ind; degeneracy_atol, gauge_atol ) - ΔA = mul!(ΔA, V * VᴴΔAV, V', 1, 1) + # VᴴΔAV is Hermitian and nonzero only in its rows and columns ind′, which are VᴴΔAVₖ' and VᴴΔAVₖ. + # For k ≤ n / 2, applying these two blocks directly, in O(n² k), is faster than forming VᴴΔAV. + ind′ = select_indices(axes(D, 1), ind) + if 2 * length(ind′) <= n + Xʳ = copy(VᴴΔAVₖ) + Xʳ[ind′, :] .= zero(eltype(Xʳ)) # these entries are part of the columns ind′ + Vₖ = view(V, :, ind′) + ΔA = mul!(ΔA, V * VᴴΔAVₖ, Vₖ', 1, 1) + ΔA = mul!(ΔA, Vₖ, Xʳ' * V', 1, 1) + else + if is_leading_index(ind′, n) # NOTE: all columns in order (e.g. `ind = Colon()`): the original path + VᴴΔAV = VᴴΔAVₖ + else + VᴴΔAV = zero!(similar(VᴴΔAVₖ, (n, n))) + VᴴΔAV[ind′, :] .= VᴴΔAVₖ' + VᴴΔAV[:, ind′] .= VᴴΔAVₖ + end + ΔA = mul!(ΔA, V * VᴴΔAV, V', 1, 1) + end return ΔA end # Diagonal: do not specialize on `A`, since we may insert `A = nothing` to assert independence of `A` in the implementation diff --git a/src/pullbacks/svd.jl b/src/pullbacks/svd.jl index 1c894528f..5245530d0 100755 --- a/src/pullbacks/svd.jl +++ b/src/pullbacks/svd.jl @@ -17,21 +17,31 @@ function check_and_prepare_svd_cotangents( indS = axes(S, 1)[ind] Δgauge = zero(eltype(S)) + # Only the columns ind′ ⊆ 1:r of UᴴΔAV are computed, its rows ind′ follow by antihermiticity. These + # are the columns of the cotangents within the rank, or all of 1:r in the full rank case + # if there are cotangents beyond it, since those have components along all of U₁ or V₁ᴴ. + j₁ = all(<=(r), indS) ? eachindex(indS) : findall(<=(r), indS) + fold = r == minmn && max(length(indU), length(indV)) > length(j₁) + ind′ = fold ? (1:r) : indS[j₁] + l₁ = fold ? indS[j₁] : eachindex(j₁) # columns of the cotangents j₁ among ind′ + k = length(ind′) + if !iszerotangent(ΔU) ΔgaugeU = zero(eltype(S)) m == size(ΔU, 1) || throw(DimensionMismatch(lazy"first dimension of ΔU ($(size(ΔU, 1))) does not match first dimension of U ($m)")) length(indU) == size(ΔU, 2) || throw(DimensionMismatch(lazy"length of selected U columns ($(length(indU))) does not match second dimension of ΔU ($(size(ΔU, 2)))")) - if indU == 1:r + if indU == ind′ ΔU₁ = copy(ΔU) else - ΔU₁ = zero(U₁) + ΔU₁ = zero!(similar(U, (m, k))) + ΔU₁[:, l₁] .= view(ΔU, :, j₁) wtmp = similar(U₁, (r,)) utmp = similar(U₁, (m,)) zeroj = Int[] for (j, i) in enumerate(indU) if i <= r - ΔU₁[:, i] .= view(ΔU, :, j) - elseif r == minmn # full rank case, ΔU₃ contains gauge-invariant information along U₁ + continue + elseif fold # full rank case, ΔU₃ contains gauge-invariant information along U₁ mul!(wtmp, U₁', view(ΔU, :, j)) mul!(ΔU₁, view(U, :, i), wtmp', -1, 1) utmp .= view(ΔU, :, j) @@ -47,27 +57,28 @@ function check_and_prepare_svd_cotangents( end UᴴΔU₁ = U₁' * ΔU₁ ΔU₊ = mul!(ΔU₁, U₁, UᴴΔU₁, -1, 1) - aUᴴΔU₁ = project_antihermitian!(UᴴΔU₁) + aUᴴΔU₁ = antihermitian_columns!(UᴴΔU₁, ind′) Δgauge = max(Δgauge, ΔgaugeU) else ΔU₊ = nothing - aUᴴΔU₁ = zero!(similar(U₁, (r, r))) + aUᴴΔU₁ = zero!(similar(U₁, (r, k))) end if !iszerotangent(ΔVᴴ) ΔgaugeV = zero(eltype(S)) n == size(ΔVᴴ, 2) || throw(DimensionMismatch(lazy"second dimension of ΔVᴴ ($(size(ΔVᴴ, 2))) does not match second dimension of Vᴴ ($n)")) length(indV) == size(ΔVᴴ, 1) || throw(DimensionMismatch(lazy"length of selected Vᴴ rows ($(length(indV))) does not match first dimension of ΔVᴴ ($(size(ΔVᴴ, 1)))")) - if indV == 1:r + if indV == ind′ ΔV₁ᴴ = copy(ΔVᴴ) else - ΔV₁ᴴ = zero(V₁ᴴ) + ΔV₁ᴴ = zero!(similar(Vᴴ, (k, n))) + ΔV₁ᴴ[l₁, :] .= view(ΔVᴴ, j₁, :) wtmp = similar(V₁ᴴ, (1, r)) vtmp = similar(V₁ᴴ, (1, n)) zeroj = Int[] for (j, i) in enumerate(indV) if i <= r - ΔV₁ᴴ[i, :] .= view(ΔVᴴ, j, :) - elseif r == minmn # full rank case, ΔV₃ contains gauge-invariant information along Vᴴ₁ + continue + elseif fold # full rank case, ΔV₃ contains gauge-invariant information along Vᴴ₁ mul!(wtmp, view(ΔVᴴ, j:j, :), V₁ᴴ') mul!(ΔV₁ᴴ, wtmp', view(Vᴴ, i:i, :), -1, 1) vtmp .= view(ΔVᴴ, j:j, :) @@ -81,41 +92,41 @@ function check_and_prepare_svd_cotangents( end VᴴΔV₁ = V₁ᴴ * ΔV₁ᴴ' ΔV₊ᴴ = mul!(ΔV₁ᴴ, VᴴΔV₁', V₁ᴴ, -1, 1) - aVᴴΔV₁ = project_antihermitian!(VᴴΔV₁) + aVᴴΔV₁ = antihermitian_columns!(VᴴΔV₁, ind′) Δgauge = max(Δgauge, ΔgaugeV) else ΔV₊ᴴ = nothing - aVᴴΔV₁ = zero!(similar(V₁ᴴ, (r, r))) + aVᴴΔV₁ = zero!(similar(V₁ᴴ, (r, k))) end - bc = Base.broadcasted(S₁', S₁, aUᴴΔU₁, aVᴴΔV₁) do s₁, s₂, u, v + Sₖ = S[ind′] + bc = Base.broadcasted(transpose(Sₖ), S₁, aUᴴΔU₁, aVᴴΔV₁) do s₁, s₂, u, v return abs(s₁ - s₂) < degeneracy_atol ? u + v : zero(u) + zero(v) end - Δgauge = max(Δgauge, maximum(abs, Base.Broadcast.instantiate(bc))) + Δgauge = max(Δgauge, maximum(abs, Base.Broadcast.instantiate(bc); init = abs(zero(eltype(S))))) if !iszerotangent(ΔSmat) ΔS = diagview(ΔSmat) length(indS) == length(ΔS) || throw(DimensionMismatch(lazy"length of selected S values ($(length(indS))) does not match length of ΔS ($(length(ΔS)))")) - bad_indS = _ind_intersect((r + 1):length(ΔS), indS) - good_indS = _ind_intersect(1:r, indS) - ΔS₁ = zero(S₁) - ΔS₁[1:length(good_indS)] .= real.(ΔS[good_indS]) - badΔS₁ = view(ΔS, bad_indS) - Δgauge = max(Δgauge, maximum(abs, badΔS₁; init = abs(zero(eltype(ΔS))))) - else - ΔS₁ = nothing + badΔS = view(ΔS, findall(>(r), indS)) + Δgauge = max(Δgauge, maximum(abs, badΔS; init = abs(zero(eltype(ΔS))))) end Δgauge ≤ gauge_atol || @warn "`svd` cotangents sensitive to gauge choice: (|Δgauge| = $Δgauge)" - UᴴΔAV = (aUᴴΔU₁ .+ aVᴴΔV₁) .* inv_safe.(S₁' .- S₁, degeneracy_atol) .+ - (aUᴴΔU₁ .- aVᴴΔV₁) .* inv_safe.(S₁' .+ S₁, degeneracy_atol) - if !iszerotangent(ΔS₁) - diagview(UᴴΔAV) .+= real.(ΔS₁) + # columns ind′ of UᴴΔAV, and the adjoint of its rows ind′ (not needed if ind′ contains all rows) + # NOTE: for all columns (k == r) only UᴴΔAV is computed, as in the original full-matrix path + UᴴΔAV = (aUᴴΔU₁ .+ aVᴴΔV₁) .* inv_safe.(transpose(Sₖ) .- S₁, degeneracy_atol) .+ + (aUᴴΔU₁ .- aVᴴΔV₁) .* inv_safe.(transpose(Sₖ) .+ S₁, degeneracy_atol) + UᴴΔAVʳ = k == r ? nothing : + (aUᴴΔU₁ .+ aVᴴΔV₁) .* inv_safe.(transpose(Sₖ) .- S₁, degeneracy_atol) .- + (aUᴴΔU₁ .- aVᴴΔV₁) .* inv_safe.(transpose(Sₖ) .+ S₁, degeneracy_atol) + if !iszerotangent(ΔSmat) + UᴴΔAV[indS[j₁] .+ r .* (l₁ .- 1)] .+= real.(view(ΔS, j₁)) # the entries (ind′[l₁], l₁) end - return UᴴΔAV, ΔU₊, ΔV₊ᴴ + return UᴴΔAV, ΔU₊, ΔV₊ᴴ, UᴴΔAVʳ, ind′ end """ @@ -161,19 +172,35 @@ function svd_pullback!( S₁ = view(S, 1:r) ΔU, ΔSmat, ΔVᴴ = ΔUSVᴴ - UᴴΔAV, ΔU₊, ΔV₊ᴴ = check_and_prepare_svd_cotangents( + UᴴΔAVₖ, ΔU₊, ΔV₊ᴴ, UᴴΔAVʳ, ind′ = check_and_prepare_svd_cotangents( U, S, Vᴴ, ΔU, ΔSmat, ΔVᴴ, r, ind; degeneracy_atol, gauge_atol ) - ΔA = mul!(ΔA, U₁, UᴴΔAV * V₁ᴴ, 1, 1) # add the contribution to ΔA + + # UᴴΔAV is nonzero only in its rows and columns ind′, which are UᴴΔAVₖ and UᴴΔAVʳ'. For k ≤ r / 2, + # applying these two blocks directly, in O(m n k), is faster than forming UᴴΔAV. + if 2 * length(ind′) <= r + UᴴΔAVʳ[ind′, :] .= zero(eltype(UᴴΔAVʳ)) # these entries are part of the columns ind′ + ΔA = mul!(ΔA, U₁ * UᴴΔAVₖ, Vᴴ[ind′, :], 1, 1) + ΔA = mul!(ΔA, U[:, ind′], UᴴΔAVʳ' * V₁ᴴ, 1, 1) + else + if is_leading_index(ind′, r) # NOTE: all columns in order (e.g. `ind = Colon()`): the original path + UᴴΔAV = UᴴΔAVₖ + else + UᴴΔAV = zero!(similar(UᴴΔAVₖ, (r, r))) + isnothing(UᴴΔAVʳ) || (UᴴΔAV[ind′, :] .= UᴴΔAVʳ') + UᴴΔAV[:, ind′] .= UᴴΔAVₖ # after the rows: only the columns hold ΔS + end + ΔA = mul!(ΔA, U₁, UᴴΔAV * V₁ᴴ, 1, 1) # add the contribution to ΔA + end # Add the remaining contributions if m > r && !iszerotangent(ΔU₊) # ΔU₁ is already orthogonal to U₁ - ΔU₊ ./= S₁' - ΔA = mul!(ΔA, ΔU₊, V₁ᴴ, 1, 1) + ΔU₊ ./= transpose(S[ind′]) + ΔA = mul!(ΔA, ΔU₊, Vᴴ[ind′, :], 1, 1) end if n > r && !iszerotangent(ΔV₊ᴴ) # ΔV₁ᴴ is already orthogonal to V₁ᴴ - ΔV₊ᴴ .= S₁ .\ ΔV₊ᴴ - ΔA = mul!(ΔA, U₁, ΔV₊ᴴ, 1, 1) + ΔV₊ᴴ .= S[ind′] .\ ΔV₊ᴴ + ΔA = mul!(ΔA, U[:, ind′], ΔV₊ᴴ, 1, 1) end return ΔA end