diff --git a/src/common/pullbacks.jl b/src/common/pullbacks.jl index 053db5335..cf17a5851 100644 --- a/src/common/pullbacks.jl +++ b/src/common/pullbacks.jl @@ -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 diff --git a/src/pullbacks/eigh.jl b/src/pullbacks/eigh.jl index 713c04b5c..12de1452f 100755 --- a/src/pullbacks/eigh.jl +++ b/src/pullbacks/eigh.jl @@ -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 diff --git a/src/pullbacks/svd.jl b/src/pullbacks/svd.jl index 1c894528f..205379aa8 100755 --- a/src/pullbacks/svd.jl +++ b/src/pullbacks/svd.jl @@ -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