Skip to content
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)
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
34 changes: 6 additions & 28 deletions src/pullbacks/eig.jl
Original file line number Diff line number Diff line change
Expand Up @@ -161,36 +161,14 @@ function eig_trunc_pullback!(
Z = ViG * VᴴΔAV

# add contribution from orthogonal complement
AP = mul!(complex.(A), V * Dmat, ViG', -1, 1)
X₀ = iszerotangent(ΔV₊) ? AP' * Z : mul!(ΔV₊, AP', Z, 1, 1)
# build the adjoint of AP directly, since that is what the series is summed with
APᴴ = mul!(complex.(A'), ViG, (V * Dmat)', -1, 1)
X₀ = iszerotangent(ΔV₊) ? APᴴ * Z : mul!(ΔV₊, APᴴ, Z, 1, 1)
X₀ ./= D'
dabsmax = maximum(abs, D)
AP ./= dabsmax
D̄⁻¹ = dabsmax ./ conj.(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.(conj.(D)), degeneracy_atol, maxiter)
Z .+= X
if eltype(ΔA) <: Real
ΔAc = mul!(AP, Z, V') # recycle AP
ΔAc = mul!(APᴴ, Z, V') # recycle APᴴ
ΔA .+= real.(ΔAc)
else
ΔA = mul!(ΔA, Z, V', 1, 1)
Expand Down
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