Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 17 additions & 0 deletions src/common/pullbacks.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Comment thread
Jutho marked this conversation as resolved.
58 changes: 35 additions & 23 deletions src/pullbacks/eigh.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
93 changes: 60 additions & 33 deletions src/pullbacks/svd.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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, :)
Expand All @@ -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

"""
Expand Down Expand Up @@ -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
Expand Down
Loading