From f3d33dadad1a3892098fb62098d4100e7674abe0 Mon Sep 17 00:00:00 2001 From: leburgel Date: Thu, 1 Oct 2026 13:48:26 +0200 Subject: [PATCH 01/10] Make `svd_pullback!` and `eigh_pullback!` cost proportional to the number of cotangent columns --- src/common/pullbacks.jl | 15 +++++++ src/pullbacks/eigh.jl | 54 +++++++++++++++--------- src/pullbacks/svd.jl | 93 ++++++++++++++++++++++++++--------------- 3 files changed, 109 insertions(+), 53 deletions(-) diff --git a/src/common/pullbacks.jl b/src/common/pullbacks.jl index cf17a5851..81fe090bd 100644 --- a/src/common/pullbacks.jl +++ b/src/common/pullbacks.jl @@ -71,3 +71,18 @@ function accelerative_smith_iteration!(X, Xₙ, G, w, atol, maxiter) end return X end + +""" + antihermitian_columns!(X, K) + +Given the columns `K` of a square matrix that is nonzero only in its columns `K`, overwrite `X` +with the same columns of the antihermitian part of that matrix. +""" +function antihermitian_columns!(X, K) + # NOTE: all columns in order (e.g. `ind = Colon()`): the original in-place projection + is_leading_index(K, size(X, 1)) && return project_antihermitian!(X) + XKK = project_antihermitian!(X[K, :]) + X ./= 2 + X[K, :] .= XKK + return X +end diff --git a/src/pullbacks/eigh.jl b/src/pullbacks/eigh.jl index 12de1452f..cc96ca6cb 100755 --- a/src/pullbacks/eigh.jl +++ b/src/pullbacks/eigh.jl @@ -4,31 +4,27 @@ 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) + K = select_indices(axes(D, 1), ind) + k = length(K) 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₁, K) 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 + Dₖ = D[K] + 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)))) @@ -36,13 +32,13 @@ function check_and_prepare_eigh_cotangents( Δ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()) + VᴴAΔV[K .+ p .* (0:(k - 1))] .+= real.(ΔD) # the entries (K[l], l) else ΔD = nothing end @@ -86,11 +82,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 K, 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. + K = select_indices(axes(D, 1), ind) + if 2 * length(K) <= n + Xʳ = copy(VᴴΔAVₖ) + Xʳ[K, :] .= zero(eltype(Xʳ)) # these entries are part of the columns K + Vₖ = V[:, K] + ΔA = mul!(ΔA, V * VᴴΔAVₖ, Vₖ', 1, 1) + ΔA = mul!(ΔA, Vₖ, Xʳ' * V', 1, 1) + else + if is_leading_index(K, 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[K, :] .= VᴴΔAVₖ' + VᴴΔAV[:, K] .= 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 205379aa8..5a554add3 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 K ⊆ 1:r of UᴴΔAV are computed, its rows K 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₁ᴴ. + jK = all(<=(r), indS) ? eachindex(indS) : findall(<=(r), indS) + fold = r == minmn && max(length(indU), length(indV)) > length(jK) + K = fold ? (1:r) : indS[jK] + lK = fold ? indS[jK] : eachindex(jK) # columns of the cotangents jK among K + k = length(K) + 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 == K ΔU₁ = copy(ΔU) else - ΔU₁ = zero(U₁) + ΔU₁ = zero!(similar(U, (m, k))) + ΔU₁[:, lK] .= view(ΔU, :, jK) 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₁, K) Δ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 == K ΔV₁ᴴ = copy(ΔVᴴ) else - ΔV₁ᴴ = zero(V₁ᴴ) + ΔV₁ᴴ = zero!(similar(Vᴴ, (k, n))) + ΔV₁ᴴ[lK, :] .= view(ΔVᴴ, jK, :) 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₁, K) Δ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[K] + 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 K of UᴴΔAV, and the adjoint of its rows K (not needed if K 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[jK] .+ r .* (lK .- 1)] .+= real.(view(ΔS, jK)) # the entries (K[lK], lK) end - return UᴴΔAV, ΔU₊, ΔV₊ᴴ + return UᴴΔAV, ΔU₊, ΔV₊ᴴ, UᴴΔAVʳ, K 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ʳ, K = 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 K, 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(K) <= r + UᴴΔAVʳ[K, :] .= zero(eltype(UᴴΔAVʳ)) # these entries are part of the columns K + ΔA = mul!(ΔA, U₁ * UᴴΔAVₖ, Vᴴ[K, :], 1, 1) + ΔA = mul!(ΔA, U[:, K], UᴴΔAVʳ' * V₁ᴴ, 1, 1) + else + if is_leading_index(K, 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[K, :] .= UᴴΔAVʳ') + UᴴΔAV[:, K] .= 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[K]) + ΔA = mul!(ΔA, ΔU₊, Vᴴ[K, :], 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[K] .\ ΔV₊ᴴ + ΔA = mul!(ΔA, U[:, K], ΔV₊ᴴ, 1, 1) end return ΔA end From 1c0255f4162cd7aeed965007807b67eaf994bdd2 Mon Sep 17 00:00:00 2001 From: leburgel Date: Fri, 2 Oct 2026 09:19:47 +0200 Subject: [PATCH 02/10] Use `ind` variable name and clarify what is actually happening --- src/common/pullbacks.jl | 18 ++++++++++-------- 1 file changed, 10 insertions(+), 8 deletions(-) diff --git a/src/common/pullbacks.jl b/src/common/pullbacks.jl index 81fe090bd..c99f03859 100644 --- a/src/common/pullbacks.jl +++ b/src/common/pullbacks.jl @@ -73,16 +73,18 @@ function accelerative_smith_iteration!(X, Xₙ, G, w, atol, maxiter) end """ - antihermitian_columns!(X, K) + antihermitian_columns!(X, ind) -Given the columns `K` of a square matrix that is nonzero only in its columns `K`, overwrite `X` -with the same columns of the antihermitian part of that matrix. +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, K) - # NOTE: all columns in order (e.g. `ind = Colon()`): the original in-place projection - is_leading_index(K, size(X, 1)) && return project_antihermitian!(X) - XKK = project_antihermitian!(X[K, :]) +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[K, :] .= XKK + X[ind, :] .= Xdiag return X end From 4ce315d77fc24cd32f00aa2fcc94b02a5a2d74b9 Mon Sep 17 00:00:00 2001 From: leburgel Date: Fri, 2 Oct 2026 09:21:00 +0200 Subject: [PATCH 03/10] Use consistent names for indices and parts withing the rank --- src/pullbacks/eigh.jl | 26 +++++++++---------- src/pullbacks/svd.jl | 58 +++++++++++++++++++++---------------------- 2 files changed, 42 insertions(+), 42 deletions(-) diff --git a/src/pullbacks/eigh.jl b/src/pullbacks/eigh.jl index cc96ca6cb..54f6fb2ce 100755 --- a/src/pullbacks/eigh.jl +++ b/src/pullbacks/eigh.jl @@ -6,8 +6,8 @@ function check_and_prepare_eigh_cotangents( # only the columns ind of VᴴΔV and VᴴAΔV are computed; their rows ind follow by antihermiticity n, p = size(V) - K = select_indices(axes(D, 1), ind) - k = length(K) + ind′ = select_indices(axes(D, 1), ind) + k = length(ind′) if !iszerotangent(ΔV) n == size(ΔV, 1) || throw(DimensionMismatch()) k == size(ΔV, 2) || throw(DimensionMismatch()) @@ -17,13 +17,13 @@ function check_and_prepare_eigh_cotangents( else ΔV₊ = mul!(copy(ΔV), V, VᴴΔV₁, -1, 1) end - aVᴴΔV₁ = antihermitian_columns!(VᴴΔV₁, K) + aVᴴΔV₁ = antihermitian_columns!(VᴴΔV₁, ind′) else ΔV₊ = nothing aVᴴΔV₁ = zero!(similar(V, (p, k))) end - Dₖ = D[K] + Dₖ = D[ind′] bc = Base.broadcasted(transpose(Dₖ), D, aVᴴΔV₁) do d₁, d₂, v return abs(d₁ - d₂) < degeneracy_atol ? v : zero(v) end @@ -38,7 +38,7 @@ function check_and_prepare_eigh_cotangents( if !iszerotangent(ΔDmat) ΔD = diagview(ΔDmat) k == length(ΔD) || throw(DimensionMismatch()) - VᴴAΔV[K .+ p .* (0:(k - 1))] .+= real.(ΔD) # the entries (K[l], l) + VᴴAΔV[ind′ .+ p .* (0:(k - 1))] .+= real.(ΔD) # the entries (ind′[l], l) else ΔD = nothing end @@ -86,22 +86,22 @@ function eigh_pullback!( D, V, ΔDmat, ΔV, ind; degeneracy_atol, gauge_atol ) - # VᴴΔAV is Hermitian and nonzero only in its rows and columns K, which are VᴴΔAVₖ' and VᴴΔAVₖ. + # 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. - K = select_indices(axes(D, 1), ind) - if 2 * length(K) <= n + ind′ = select_indices(axes(D, 1), ind) + if 2 * length(ind′) <= n Xʳ = copy(VᴴΔAVₖ) - Xʳ[K, :] .= zero(eltype(Xʳ)) # these entries are part of the columns K - Vₖ = V[:, K] + Xʳ[ind′, :] .= zero(eltype(Xʳ)) # these entries are part of the columns ind′ + Vₖ = V[:, ind′] ΔA = mul!(ΔA, V * VᴴΔAVₖ, Vₖ', 1, 1) ΔA = mul!(ΔA, Vₖ, Xʳ' * V', 1, 1) else - if is_leading_index(K, n) # NOTE: all columns in order (e.g. `ind = Colon()`): the original path + 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[K, :] .= VᴴΔAVₖ' - VᴴΔAV[:, K] .= VᴴΔAVₖ + VᴴΔAV[ind′, :] .= VᴴΔAVₖ' + VᴴΔAV[:, ind′] .= VᴴΔAVₖ end ΔA = mul!(ΔA, V * VᴴΔAV, V', 1, 1) end diff --git a/src/pullbacks/svd.jl b/src/pullbacks/svd.jl index 5a554add3..966cec9e1 100755 --- a/src/pullbacks/svd.jl +++ b/src/pullbacks/svd.jl @@ -17,24 +17,24 @@ function check_and_prepare_svd_cotangents( indS = axes(S, 1)[ind] Δgauge = zero(eltype(S)) - # Only the columns K ⊆ 1:r of UᴴΔAV are computed, its rows K follow by antihermiticity. These + # 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₁ᴴ. - jK = all(<=(r), indS) ? eachindex(indS) : findall(<=(r), indS) - fold = r == minmn && max(length(indU), length(indV)) > length(jK) - K = fold ? (1:r) : indS[jK] - lK = fold ? indS[jK] : eachindex(jK) # columns of the cotangents jK among K - k = length(K) + 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 == K + if indU == ind′ ΔU₁ = copy(ΔU) else ΔU₁ = zero!(similar(U, (m, k))) - ΔU₁[:, lK] .= view(ΔU, :, jK) + ΔU₁[:, l₁] .= view(ΔU, :, j₁) wtmp = similar(U₁, (r,)) utmp = similar(U₁, (m,)) zeroj = Int[] @@ -57,7 +57,7 @@ function check_and_prepare_svd_cotangents( end UᴴΔU₁ = U₁' * ΔU₁ ΔU₊ = mul!(ΔU₁, U₁, UᴴΔU₁, -1, 1) - aUᴴΔU₁ = antihermitian_columns!(UᴴΔU₁, K) + aUᴴΔU₁ = antihermitian_columns!(UᴴΔU₁, ind′) Δgauge = max(Δgauge, ΔgaugeU) else ΔU₊ = nothing @@ -67,11 +67,11 @@ function check_and_prepare_svd_cotangents( Δ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 == K + if indV == ind′ ΔV₁ᴴ = copy(ΔVᴴ) else ΔV₁ᴴ = zero!(similar(Vᴴ, (k, n))) - ΔV₁ᴴ[lK, :] .= view(ΔVᴴ, jK, :) + ΔV₁ᴴ[l₁, :] .= view(ΔVᴴ, j₁, :) wtmp = similar(V₁ᴴ, (1, r)) vtmp = similar(V₁ᴴ, (1, n)) zeroj = Int[] @@ -92,14 +92,14 @@ function check_and_prepare_svd_cotangents( end VᴴΔV₁ = V₁ᴴ * ΔV₁ᴴ' ΔV₊ᴴ = mul!(ΔV₁ᴴ, VᴴΔV₁', V₁ᴴ, -1, 1) - aVᴴΔV₁ = antihermitian_columns!(VᴴΔV₁, K) + aVᴴΔV₁ = antihermitian_columns!(VᴴΔV₁, ind′) Δgauge = max(Δgauge, ΔgaugeV) else ΔV₊ᴴ = nothing aVᴴΔV₁ = zero!(similar(V₁ᴴ, (r, k))) end - Sₖ = S[K] + 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 @@ -115,7 +115,7 @@ function check_and_prepare_svd_cotangents( Δgauge ≤ gauge_atol || @warn "`svd` cotangents sensitive to gauge choice: (|Δgauge| = $Δgauge)" - # columns K of UᴴΔAV, and the adjoint of its rows K (not needed if K contains all rows) + # 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) @@ -123,10 +123,10 @@ function check_and_prepare_svd_cotangents( (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[jK] .+ r .* (lK .- 1)] .+= real.(view(ΔS, jK)) # the entries (K[lK], lK) + UᴴΔAV[indS[j₁] .+ r .* (l₁ .- 1)] .+= real.(view(ΔS, j₁)) # the entries (ind′[l₁], l₁) end - return UᴴΔAV, ΔU₊, ΔV₊ᴴ, UᴴΔAVʳ, K + return UᴴΔAV, ΔU₊, ΔV₊ᴴ, UᴴΔAVʳ, ind′ end """ @@ -172,35 +172,35 @@ function svd_pullback!( S₁ = view(S, 1:r) ΔU, ΔSmat, ΔVᴴ = ΔUSVᴴ - UᴴΔAVₖ, ΔU₊, ΔV₊ᴴ, UᴴΔAVʳ, K = 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 ) - # UᴴΔAV is nonzero only in its rows and columns K, which are UᴴΔAVₖ and UᴴΔAVʳ'. For k ≤ r / 2, + # 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(K) <= r - UᴴΔAVʳ[K, :] .= zero(eltype(UᴴΔAVʳ)) # these entries are part of the columns K - ΔA = mul!(ΔA, U₁ * UᴴΔAVₖ, Vᴴ[K, :], 1, 1) - ΔA = mul!(ΔA, U[:, K], UᴴΔAVʳ' * V₁ᴴ, 1, 1) + 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(K, r) # NOTE: all columns in order (e.g. `ind = Colon()`): the original path + 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[K, :] .= UᴴΔAVʳ') - UᴴΔAV[:, K] .= UᴴΔAVₖ # after the rows: only the columns hold ΔS + 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₊ ./= transpose(S[K]) - ΔA = mul!(ΔA, ΔU₊, Vᴴ[K, :], 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[K] .\ ΔV₊ᴴ - ΔA = mul!(ΔA, U[:, K], ΔV₊ᴴ, 1, 1) + ΔV₊ᴴ .= S[ind′] .\ ΔV₊ᴴ + ΔA = mul!(ΔA, U[:, ind′], ΔV₊ᴴ, 1, 1) end return ΔA end From a1fb1bac934d37f9bbf6bac26395864276b18950 Mon Sep 17 00:00:00 2001 From: Lander Burgelman <39218680+leburgel@users.noreply.github.com> Date: Fri, 2 Oct 2026 16:08:17 +0200 Subject: [PATCH 04/10] Update src/pullbacks/eigh.jl Co-authored-by: Jutho --- src/pullbacks/eigh.jl | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/src/pullbacks/eigh.jl b/src/pullbacks/eigh.jl index 54f6fb2ce..02199b97e 100755 --- a/src/pullbacks/eigh.jl +++ b/src/pullbacks/eigh.jl @@ -23,11 +23,9 @@ function check_and_prepare_eigh_cotangents( aVᴴΔV₁ = zero!(similar(V, (p, k))) end - Dₖ = D[ind′] - 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)" From c12eee3d43ffba90581046c650bb8074580370ef Mon Sep 17 00:00:00 2001 From: Lander Burgelman <39218680+leburgel@users.noreply.github.com> Date: Fri, 2 Oct 2026 16:08:42 +0200 Subject: [PATCH 05/10] Update src/pullbacks/eigh.jl Co-authored-by: Jutho --- src/pullbacks/eigh.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/pullbacks/eigh.jl b/src/pullbacks/eigh.jl index 02199b97e..2f4948052 100755 --- a/src/pullbacks/eigh.jl +++ b/src/pullbacks/eigh.jl @@ -90,7 +90,7 @@ function eigh_pullback!( if 2 * length(ind′) <= n Xʳ = copy(VᴴΔAVₖ) Xʳ[ind′, :] .= zero(eltype(Xʳ)) # these entries are part of the columns ind′ - Vₖ = V[:, ind′] + Vₖ = view(V, :, ind′) ΔA = mul!(ΔA, V * VᴴΔAVₖ, Vₖ', 1, 1) ΔA = mul!(ΔA, Vₖ, Xʳ' * V', 1, 1) else From ea0dbf66a95de5e0a024b1e4963f1e65e156f307 Mon Sep 17 00:00:00 2001 From: Lander Burgelman <39218680+leburgel@users.noreply.github.com> Date: Fri, 2 Oct 2026 16:19:00 +0200 Subject: [PATCH 06/10] Update src/pullbacks/eigh.jl Co-authored-by: Jutho --- src/pullbacks/eigh.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/pullbacks/eigh.jl b/src/pullbacks/eigh.jl index 2f4948052..f89aa890d 100755 --- a/src/pullbacks/eigh.jl +++ b/src/pullbacks/eigh.jl @@ -36,7 +36,7 @@ function check_and_prepare_eigh_cotangents( if !iszerotangent(ΔDmat) ΔD = diagview(ΔDmat) k == length(ΔD) || throw(DimensionMismatch()) - VᴴAΔV[ind′ .+ p .* (0:(k - 1))] .+= real.(ΔD) # the entries (ind′[l], l) + diagview(view(VᴴAΔV, ind′, :)) .+= real.(ΔD) else ΔD = nothing end From 72549b85667eaff765a72f4202ec133f613af92e Mon Sep 17 00:00:00 2001 From: leburgel Date: Sat, 3 Oct 2026 08:32:17 +0200 Subject: [PATCH 07/10] Update src/pullbacks/eigh.jl Co-authored-by: Jutho --- src/pullbacks/eigh.jl | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/src/pullbacks/eigh.jl b/src/pullbacks/eigh.jl index f89aa890d..ee41e93cc 100755 --- a/src/pullbacks/eigh.jl +++ b/src/pullbacks/eigh.jl @@ -88,11 +88,10 @@ function eigh_pullback!( # 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) + VᴴΔAVₖ[ind′, :] .= zero(eltype(VᴴΔAVₖ)) + ΔA = mul!(ΔA, Vₖ, VᴴΔAVₖ' * 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ₖ From fc18c5620347b4b592a58420f7a996c226c55287 Mon Sep 17 00:00:00 2001 From: leburgel Date: Sat, 3 Oct 2026 08:32:17 +0200 Subject: [PATCH 08/10] Update src/pullbacks/svd.jl Co-authored-by: Jutho --- src/pullbacks/svd.jl | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/src/pullbacks/svd.jl b/src/pullbacks/svd.jl index 966cec9e1..0dd6b8f70 100755 --- a/src/pullbacks/svd.jl +++ b/src/pullbacks/svd.jl @@ -178,10 +178,13 @@ function svd_pullback!( # 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. + S′ = view(S, ind′) + U′ = view(U, :, ind′) + Vᴴ′ = view(Vᴴ, ind′, :) # this might be slightly confusion with adjoint 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) + ΔA = mul!(ΔA, U₁ * UᴴΔAVₖ, Vᴴ′, 1, 1) + ΔA = mul!(ΔA, U′, 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ₖ @@ -195,12 +198,12 @@ function svd_pullback!( # Add the remaining contributions if m > r && !iszerotangent(ΔU₊) # ΔU₁ is already orthogonal to U₁ - ΔU₊ ./= transpose(S[ind′]) - ΔA = mul!(ΔA, ΔU₊, Vᴴ[ind′, :], 1, 1) + ΔU₊ ./= transpose(S′) + ΔA = mul!(ΔA, ΔU₊, Vᴴ′, 1, 1) end if n > r && !iszerotangent(ΔV₊ᴴ) # ΔV₁ᴴ is already orthogonal to V₁ᴴ - ΔV₊ᴴ .= S[ind′] .\ ΔV₊ᴴ - ΔA = mul!(ΔA, U[:, ind′], ΔV₊ᴴ, 1, 1) + ΔV₊ᴴ .= S′ .\ ΔV₊ᴴ + ΔA = mul!(ΔA, U′, ΔV₊ᴴ, 1, 1) end return ΔA end From 151087a0186c868b3f6dcc0fd1c638bc56547285 Mon Sep 17 00:00:00 2001 From: leburgel Date: Sat, 3 Oct 2026 10:27:19 +0200 Subject: [PATCH 09/10] =?UTF-8?q?Use=20copies=20instead=20of=20views=20of?= =?UTF-8?q?=20the=20kept=20columns=20of=20U,=20V=E1=B4=B4=20and=20V?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Views with a vector index fall back to scalar indexing in mul! on GPU. --- src/pullbacks/eigh.jl | 2 +- src/pullbacks/svd.jl | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/src/pullbacks/eigh.jl b/src/pullbacks/eigh.jl index ee41e93cc..ec9abcc40 100755 --- a/src/pullbacks/eigh.jl +++ b/src/pullbacks/eigh.jl @@ -88,7 +88,7 @@ function eigh_pullback!( # 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 - Vₖ = view(V, :, ind′) + Vₖ = V[:, ind′] ΔA = mul!(ΔA, V * VᴴΔAVₖ, Vₖ', 1, 1) VᴴΔAVₖ[ind′, :] .= zero(eltype(VᴴΔAVₖ)) ΔA = mul!(ΔA, Vₖ, VᴴΔAVₖ' * V', 1, 1) diff --git a/src/pullbacks/svd.jl b/src/pullbacks/svd.jl index 0dd6b8f70..fe0847966 100755 --- a/src/pullbacks/svd.jl +++ b/src/pullbacks/svd.jl @@ -179,8 +179,8 @@ function svd_pullback!( # 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. S′ = view(S, ind′) - U′ = view(U, :, ind′) - Vᴴ′ = view(Vᴴ, ind′, :) # this might be slightly confusion with adjoint + U′ = U[:, ind′] + Vᴴ′ = Vᴴ[ind′, :] # this might be slightly confusion with adjoint 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ᴴ′, 1, 1) From 64e4a1d8d17c76f4adea5ca3da08c6adaa2b5fe0 Mon Sep 17 00:00:00 2001 From: leburgel Date: Sat, 3 Oct 2026 10:27:24 +0200 Subject: [PATCH 10/10] =?UTF-8?q?Rename=20the=20kept-column=20factors=20in?= =?UTF-8?q?=20svd=5Fpullback!=20to=20S=E2=82=96,=20U=E2=82=96=20and=20V?= =?UTF-8?q?=E1=B4=B4=E2=82=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/pullbacks/svd.jl | 18 +++++++++--------- 1 file changed, 9 insertions(+), 9 deletions(-) diff --git a/src/pullbacks/svd.jl b/src/pullbacks/svd.jl index fe0847966..8b129e356 100755 --- a/src/pullbacks/svd.jl +++ b/src/pullbacks/svd.jl @@ -178,13 +178,13 @@ function svd_pullback!( # 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. - S′ = view(S, ind′) - U′ = U[:, ind′] - Vᴴ′ = Vᴴ[ind′, :] # this might be slightly confusion with adjoint + Sₖ = view(S, ind′) + Uₖ = U[:, ind′] + Vᴴₖ = Vᴴ[ind′, :] 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ᴴ′, 1, 1) - ΔA = mul!(ΔA, U′, UᴴΔAVʳ' * V₁ᴴ, 1, 1) + ΔA = mul!(ΔA, U₁ * UᴴΔAVₖ, Vᴴₖ, 1, 1) + ΔA = mul!(ΔA, Uₖ, 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ₖ @@ -198,12 +198,12 @@ function svd_pullback!( # Add the remaining contributions if m > r && !iszerotangent(ΔU₊) # ΔU₁ is already orthogonal to U₁ - ΔU₊ ./= transpose(S′) - ΔA = mul!(ΔA, ΔU₊, Vᴴ′, 1, 1) + ΔU₊ ./= transpose(Sₖ) + ΔA = mul!(ΔA, ΔU₊, Vᴴₖ, 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ₖ .\ ΔV₊ᴴ + ΔA = mul!(ΔA, Uₖ, ΔV₊ᴴ, 1, 1) end return ΔA end