Skip to content
Merged
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
36 changes: 36 additions & 0 deletions src/common/pullbacks.jl
Original file line number Diff line number Diff line change
Expand Up @@ -35,3 +35,39 @@ 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)

"""
accelerative_smith_iteration!(X, Xₙ, G, w, atol, maxiter)

Solve `X = B + G * X * Diagonal(w)` by summing the Neumann series
`X = Σₖ Gᵏ * B * Diagonal(w)ᵏ` by doubling (Smith's method), i.e. by repeatedly adding
`G^(2ʲ) * X * Diagonal(w)^(2ʲ)` to `X` until the norm of that increment drops below `atol`,
for at most `maxiter` steps.

On entry, `X` contains `B`, and it is overwritten with the result. `Xₙ` is used as a buffer,
and `G` and `w` are overwritten. `w` is normalized such that `maximum(abs, w) == 1`, so that
squaring it can only shrink it; `G` is scaled by the inverse factor to compensate.

Reference: https://doi.org/10.1016/j.aml.2009.01.012.
"""
function accelerative_smith_iteration!(X, Xₙ, G, w, atol, maxiter)
Gₙ = similar(G)
Comment thread
leburgel marked this conversation as resolved.
wmax = maximum(abs, w)
w ./= wmax
G .*= wmax
for k in 1:maxiter
Xₙ = rmul!(mul!(Xₙ, G, X), Diagonal(w))
if maximum(abs, Xₙ) < atol
break
end
X .+= Xₙ
if k == maxiter
@warn "Sylvester iteration did not converge after $k iterations, final norm of X: $(maximum(abs, X))"
break
end
w .= w .^ 2
Gₙ = mul!(Gₙ, G, G)
G, Gₙ = Gₙ, G
end
return X
end
31 changes: 3 additions & 28 deletions src/pullbacks/eigh.jl
Original file line number Diff line number Diff line change
Expand Up @@ -153,38 +153,13 @@ function eigh_trunc_pullback!(
if !iszerotangent(ΔV₊)
X₀ = rdiv!(ΔV₊, Diagonal(D))
AP = mul!(copy(A), V * Dmat, V', -1, 1)
# Normalize by the smallest retained |eigenvalue|, as `svd_trunc_pullback!` does
# with `S[end]`. That caps `max|D⁻¹|` at 1, so squaring can only shrink it.
dabsmin = minimum(abs, D)
AP ./= dabsmin
D⁻¹ = dabsmin ./ D
X₁ = rmul!(AP * X₀, Diagonal(D⁻¹))
X₁ .+= X₀
Xₖ, Xₖ₊₁ = X₁, X₀
APₖ, APₖ₊₁ = AP * AP, AP
D⁻¹ₖ, D⁻¹ₖ₊₁ = D⁻¹ .^ 2, D⁻¹
for k in 1:maxiter
Xₖ₊₁ = rmul!(mul!(Xₖ₊₁, APₖ, Xₖ), Diagonal(D⁻¹ₖ))
if norm(Xₖ₊₁, Inf) < degeneracy_atol
break
end
Xₖ₊₁ .+= Xₖ
if k == maxiter
@warn "Sylvester iteration did not converge after $k iterations, final norm of X: $(norm(Xₖ₊₁, Inf)))"
break
end
D⁻¹ₖ₊₁ .= D⁻¹ₖ .^ 2
APₖ₊₁ = mul!(APₖ₊₁, APₖ, APₖ)
Xₖ, Xₖ₊₁ = Xₖ₊₁, Xₖ
APₖ, APₖ₊₁ = APₖ₊₁, APₖ
D⁻¹ₖ, D⁻¹ₖ₊₁ = D⁻¹ₖ₊₁, D⁻¹ₖ
end
Z .+= Xₖ
X = accelerative_smith_iteration!(X₀, similar(X₀), AP, inv.(D), degeneracy_atol, maxiter)
Z .+= X
# we cannot directly multiply Z * V' into ΔA, because we have to
# take the Hermitian part, and cannot apply project_hermitian! to
# the current contents of ΔA
# TODO: add an `add_project_hermitian!`
# recycle AP's storage, but overwrite it: the loop leaves APₖ in that buffer
# recycle AP's storage, but overwrite it: `accelerative_smith_iteration!` may leave a power of AP in it
ΔA′ = project_hermitian!(mul!(AP, Z, V'))
ΔA .+= ΔA′
else
Expand Down
53 changes: 19 additions & 34 deletions src/pullbacks/svd.jl
Original file line number Diff line number Diff line change
Expand Up @@ -247,41 +247,26 @@ function svd_trunc_pullback!(
Y₀ᴴ = iszerotangent(ΔV₊ᴴ) ? zero(Vᴴ) : ldiv!(Diagonal(S), ΔV₊ᴴ)
US = mul!(ΔAV, U, Smat) # recycle ΔAV
AP = mul!(copy(A), US, Vᴴ, -1, 1)
minS = @view S[end:end]
AP ./= minS
S⁻¹ = minS ./ S
X₁ = rmul!(AP * Y₀ᴴ', Diagonal(S⁻¹))
X₁ .+= X₀
Y₁ᴴ = lmul!(Diagonal(S⁻¹), X₀' * AP)
Y₁ᴴ .+= Y₀ᴴ
Xₖ, Xₖ₊₁ = X₁, X₀
Yₖᴴ, Yₖ₊₁ᴴ = Y₁ᴴ, Y₀ᴴ
APAᴴₖ, AᴴPAₖ = AP * AP', AP' * AP
APAᴴₖ₊₁, AᴴPAₖ₊₁ = zero(APAᴴₖ), zero(AᴴPAₖ)
S⁻¹ₖ, S⁻¹ₖ₊₁ = S⁻¹ .^ 2, S⁻¹
for k in 1:maxiter
Xₖ₊₁ = rmul!(mul!(Xₖ₊₁, APAᴴₖ, Xₖ), Diagonal(S⁻¹ₖ))
Yₖ₊₁ᴴ = lmul!(Diagonal(S⁻¹ₖ), mul!(Yₖ₊₁ᴴ, Yₖᴴ, AᴴPAₖ))
if norm(Xₖ₊₁, Inf) < degeneracy_atol && norm(Yₖ₊₁ᴴ, Inf) < degeneracy_atol
break
end
Xₖ₊₁ .+= Xₖ
Yₖ₊₁ᴴ .+= Yₖᴴ
if k == maxiter
@warn "Sylvester iteration did not converge after $k iterations, final norms of X: $(norm(Xₖ₊₁, Inf)), Yᴴ: $(norm(Yₖ₊₁ᴴ, Inf)))"
break
end
S⁻¹ₖ₊₁ .= S⁻¹ₖ .^ 2
APAᴴₖ₊₁ = mul!(APAᴴₖ₊₁, APAᴴₖ, APAᴴₖ)
AᴴPAₖ₊₁ = mul!(AᴴPAₖ₊₁, AᴴPAₖ, AᴴPAₖ)
Xₖ, Xₖ₊₁ = Xₖ₊₁, Xₖ
Yₖᴴ, Yₖ₊₁ᴴ = Yₖ₊₁ᴴ, Yₖᴴ
APAᴴₖ, APAᴴₖ₊₁ = APAᴴₖ₊₁, APAᴴₖ
AᴴPAₖ, AᴴPAₖ₊₁ = AᴴPAₖ₊₁, AᴴPAₖ
S⁻¹ₖ, S⁻¹ₖ₊₁ = S⁻¹ₖ₊₁, S⁻¹ₖ
S⁻¹ = inv.(S)
# sum the series on the smaller side only, the other side follows from
# Yᴴ = Y₀ᴴ + S⁻¹ X' AP (m ≤ n) or X = X₀ + AP Y S⁻¹ (m > n)
if m ≤ n
X = rmul!(AP * Y₀ᴴ', Diagonal(S⁻¹))
X .+= X₀
X = accelerative_smith_iteration!(X, X₀, AP * AP', S⁻¹ .^ 2, degeneracy_atol, maxiter) # recycle X₀
Yᴴ = lmul!(Diagonal(S⁻¹), X' * AP)
Yᴴ .+= Y₀ᴴ
ΔA = mul!(ΔA, X, Vᴴ, 1, 1)
ΔA = mul!(ΔA, U, Yᴴ, 1, 1)
else
Y = rmul!(AP' * X₀, Diagonal(S⁻¹))
Y .+= Y₀ᴴ'
Y = accelerative_smith_iteration!(Y, similar(Y), AP' * AP, S⁻¹ .^ 2, degeneracy_atol, maxiter)
X = rmul!(AP * Y, Diagonal(S⁻¹))
X .+= X₀
ΔA = mul!(ΔA, X, Vᴴ, 1, 1)
ΔA = mul!(ΔA, U, Y', 1, 1)
end
ΔA = mul!(ΔA, Xₖ, Vᴴ, 1, 1)
ΔA = mul!(ΔA, U, Yₖᴴ, 1, 1)
end
return ΔA
end
Expand Down
Loading