diff --git a/src/MatrixAlgebraKit.jl b/src/MatrixAlgebraKit.jl index 78d36bfe5..9fa6f6644 100644 --- a/src/MatrixAlgebraKit.jl +++ b/src/MatrixAlgebraKit.jl @@ -92,6 +92,7 @@ include("common/defaults.jl") include("common/householder.jl") include("common/initialization.jl") include("common/pullbacks.jl") +include("common/stein.jl") include("common/safemethods.jl") include("common/view.jl") include("common/regularinv.jl") diff --git a/src/common/pullbacks.jl b/src/common/pullbacks.jl index c99f03859..43bba4c45 100644 --- a/src/common/pullbacks.jl +++ b/src/common/pullbacks.jl @@ -36,42 +36,6 @@ 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 - """ antihermitian_columns!(X, ind) diff --git a/src/common/stein.jl b/src/common/stein.jl new file mode 100644 index 000000000..e0b579a7c --- /dev/null +++ b/src/common/stein.jl @@ -0,0 +1,177 @@ +# Solvers for the Stein equation X - G X Diagonal(w) = B, i.e. (1 - wᵢ G) xᵢ = bᵢ per column. +# Used in the pullback of truncated decompositions, with G = P Pᴴ or Pᴴ P, wᵢ = 1/σᵢ² for the case of SVD, +# and G = P, wᵢ = 1/λᵢ for the case of EIG(H). Here, P the part of A outside the kept vectors. +# Naive iteration requires γᵢ, the spectral radius of wᵢ G, to be smaller than 1. + +""" + 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 + +# smallest Ritz value: smallest eigenvalue of the Lanczos tridiagonal from the CG coefficients `α`, `β` +function _cg_ritz_min(α, β) + l = length(α) + d = similar(α) + d[1] = 1 / α[1] + for j in 2:l + d[j] = 1 / α[j] + β[j - 1] / α[j - 1] + end + e = sqrt.(β[1:(l - 1)]) ./ α[1:(l - 1)] + return LinearAlgebra.eigmin(LinearAlgebra.SymTridiagonal(d, e)) +end + +# (1 - wᵢ G) applied to the columns of Z +_stein_op(G, w, Z) = Z .- (G * Z) .* transpose(w) + +# iterations until the squared residual norm `r` drops below `tol²`, at the faster of its rates over +# the last `l` iterations (from `r₀`) and all `m` (from `rᵢ`) +function _cg_remaining(r, r₀, rᵢ, tol, l, m) + q = min(log(r / r₀) / l, log(r / rᵢ) / m) # logarithm of the reduction per iteration + return q < 0 ? log(tol^2 / r) / q : oftype(float(r), Inf) +end + +# real parts of the column-wise inner products of A and B, on the device of A and B. CPU arrays use +# `dot` per column, which avoids the temporary. This is restricted to `Matrix` rather than +# `StridedMatrix`, which GPU arrays also are: there it would return a CPU vector, one `dot` call +# each, which the GPU broadcasts in `hermitian_stein_cg!` cannot mix with the device arrays. +_coldots(A, B) = vec(real(sum(conj.(A) .* B; dims = 1))) +const _CPUMatrix = Union{Matrix, SubArray{<:Any, 2, <:Matrix}} +_coldots(A::_CPUMatrix, B::_CPUMatrix) = [real(LinearAlgebra.dot(view(A, :, j), view(B, :, j))) for j in axes(A, 2)] + +""" + hermitian_stein_cg!(X, applyG!, formG, w, atol, maxiter; cost_apply, cost_apply_formed, cost_form, cost_square = nothing) + +Solve `X - G * X * Diagonal(w) = B` for Hermitian `G` by conjugate gradients on all columns at +once, in at most `maxiter` iterations: O(n² k) per iteration for k columns, against O(n³) per +doubling step. Only the products wᵢ G enter, so neither needs normalizing. `X` contains `B` on +entry and is overwritten with the result. + +`applyG!(Y, Z)` sets `Y = G * Z` without forming `G`, and `formG()` returns `G`. `G` is formed +once the predicted remaining applications save more than `cost_form`, given the costs per column +`cost_apply` (by `applyG!`) and `cost_apply_formed` (by `G`); `cost_form = 0` forms it at once. +With `cost_square`, for indefinite `G` (eigh), the solver restarts on the residual equation +multiplied by 1 + wᵢ G, (1 - wᵢ² G²) xᵢ = (1 + wᵢ G) bᵢ, once half the predicted remaining cost +exceeds `cost_square` (forming G² and the new right-hand side): its spectrum [1 - γᵢ², 1] needs +about half the iterations. A column stops once its residual is below `atol` times the smallest +Ritz value of the slowest column, so that its error is below about `atol`. +""" +function hermitian_stein_cg!( + X, applyG!, formG, w, atol, maxiter; + cost_apply, cost_apply_formed = cost_apply, cost_form, cost_square = nothing, nprobe::Int = 5 + ) + RT = real(eltype(X)) + G = iszero(cost_form) ? formG() : nothing + tol = RT(atol) # stopping residual 2-norm, refined by the Ritz values + ρ = _coldots(X, X) + cols = findall(ρ .> tol^2) # active columns, kept contiguous at the front + nact = length(cols) + iszero(nact) && return fill!(X, zero(eltype(X))) + R = X[:, cols] # X keeps B until the end + P = zero(R) + Q = similar(R) + Xc = zero(R) + wc = w[cols] + ρ = ρ[cols] + β = zero(ρ) + ρ₀ = copy(ρ) # squared residual norms at the previous check + ρᵢ = copy(ρ) # and at the start + αs = [RT[] for _ in 1:nact] + βs = [RT[] for _ in 1:nact] + for numiter in 1:maxiter + if numiter > 1 && (numiter - 1) % nprobe == 0 + c = argmax(abs.(view(wc, 1:nact))) # the slowest column + tol = atol * _cg_ritz_min(αs[c], βs[c]) + nact = _cg_compact!( + view(ρ, 1:nact), tol, view(R, :, 1:nact), view(P, :, 1:nact), view(Xc, :, 1:nact), view(wc, 1:nact), + view(β, 1:nact), view(ρ₀, 1:nact), view(ρᵢ, 1:nact), view(cols, 1:nact), view(αs, 1:nact), view(βs, 1:nact) + ) + iszero(nact) && break + # predicted column applications + napply = sum(_cg_remaining.(view(ρ, 1:nact), view(ρ₀, 1:nact), view(ρᵢ, 1:nact), tol, nprobe, numiter - 1)) + if isnothing(G) && napply * (cost_apply - cost_apply_formed) > cost_form + G = formG() + end + if !isnothing(cost_square) && napply * cost_apply / 2 > cost_square # restart on the squared equation + isnothing(G) && (G = formG()) + X₁ = zero(X) # the current solution + X₁[:, cols] .= Xc + R = X .- _stein_op(G, w, X₁) + X .= R .+ (G * R) .* transpose(w) + G² = G * G + hermitian_stein_cg!(X, nothing, () -> G², w .^ 2, atol, maxiter; cost_apply, cost_form = 0, nprobe) + return X .+= X₁ + end + ρ₀ .= ρ + end + Pₐ, Rₐ, Qₐ, wₐ = view(P, :, 1:nact), view(R, :, 1:nact), view(Q, :, 1:nact), view(wc, 1:nact) + ρₐ, βₐ = view(ρ, 1:nact), view(β, 1:nact) + Pₐ .= Rₐ .+ Pₐ .* transpose(βₐ) + isnothing(G) ? applyG!(Qₐ, Pₐ) : mul!(Qₐ, G, Pₐ) + Qₐ .= Pₐ .- Qₐ .* transpose(wₐ) # q = (1 - wᵢ G) p + α = ρₐ ./ _coldots(Pₐ, Qₐ) + view(Xc, :, 1:nact) .+= Pₐ .* transpose(α) + Rₐ .-= Qₐ .* transpose(α) + βₐ .= ρₐ # ρold + ρₐ .= _coldots(Rₐ, Rₐ) + βₐ .= ρₐ ./ βₐ + for (j, a, b) in zip(1:nact, Array(α), Array(βₐ)) + push!(αs[j], a) + push!(βs[j], b) + end + nact = _cg_compact!( + ρₐ, tol, Rₐ, Pₐ, view(Xc, :, 1:nact), wₐ, + βₐ, view(ρ₀, 1:nact), view(ρᵢ, 1:nact), view(cols, 1:nact), view(αs, 1:nact), view(βs, 1:nact) + ) + iszero(nact) && break + end + iszero(nact) || @warn "conjugate gradients did not converge in $maxiter iterations, largest residual norm: $(sqrt(maximum(view(ρ, 1:nact))))" + fill!(X, zero(eltype(X))) + X[:, cols] .= Xc + return X +end + +# move the columns with `ρ > tol²` to the front of all arrays (in order) and return their number +function _cg_compact!(ρ, tol, arrays...) + keep = Array(ρ) .> tol^2 + all(keep) && return length(keep) + perm = vcat(findall(keep), findall(.!keep)) + for a in (ρ, arrays...) + if a isa AbstractMatrix + a .= a[:, perm] + else + a .= a[perm] + end + end + return count(keep) +end diff --git a/src/pullbacks/eigh.jl b/src/pullbacks/eigh.jl index 2b3afacd0..6f8063f58 100755 --- a/src/pullbacks/eigh.jl +++ b/src/pullbacks/eigh.jl @@ -143,7 +143,7 @@ function eigh_trunc_pullback!( ΔA::AbstractMatrix, A, DV, ΔDV; degeneracy_atol::Real = default_pullback_rank_atol(DV[1]), gauge_atol::Real = default_pullback_gauge_atol(ΔDV[2]), - maxiter::Int = 100 # TODO: better default, depending on expected number of steps using quadratic convergence? + maxiter::Int = 10 * size(ΔA, 1) # conjugate-gradient iterations ) # Basic size checks and determination @@ -162,13 +162,16 @@ function eigh_trunc_pullback!( if !iszerotangent(ΔV₊) X₀ = rdiv!(ΔV₊, Diagonal(D)) AP = mul!(copy(A), V * Dmat, V', -1, 1) - X = accelerative_smith_iteration!(X₀, similar(X₀), AP, inv.(D), degeneracy_atol, maxiter) + X = hermitian_stein_cg!( + X₀, nothing, () -> AP, inv.(D), degeneracy_atol, maxiter; + cost_apply = n^2, cost_form = 0, cost_square = n^3 + n^2 * p + ) 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: `accelerative_smith_iteration!` may leave a power of AP in it + # recycle AP's storage ΔA′ = project_hermitian!(mul!(AP, Z, V')) ΔA .+= ΔA′ else diff --git a/src/pullbacks/svd.jl b/src/pullbacks/svd.jl index 576997b69..8fa8904dc 100755 --- a/src/pullbacks/svd.jl +++ b/src/pullbacks/svd.jl @@ -223,7 +223,7 @@ function svd_trunc_pullback!( rank_atol::Real = 0, degeneracy_atol::Real = default_pullback_rank_atol(USVᴴ[2]), gauge_atol::Real = default_pullback_gauge_atol(ΔUSVᴴ...), - maxiter::Int = 100 # TODO: better default, depending on expected number of steps using quadratic convergence? + maxiter::Int = 10 * minimum(size(ΔA)) # conjugate-gradient iterations ) # Extract the SVD components U, Smat, Vᴴ = USVᴴ @@ -257,7 +257,12 @@ function svd_trunc_pullback!( 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₀ + APᴴZ = similar(X, n, p) # for applying AP AP' without forming it + X = hermitian_stein_cg!( + X, (GZ, Z) -> mul!(GZ, AP, mul!(view(APᴴZ, :, axes(Z, 2)), AP', Z)), () -> AP * AP', + S⁻¹ .^ 2, degeneracy_atol, maxiter; + cost_apply = 2 * m * n, cost_apply_formed = m^2, cost_form = m^2 * n + ) Yᴴ = lmul!(Diagonal(S⁻¹), X' * AP) Yᴴ .+= Y₀ᴴ ΔA = mul!(ΔA, X, Vᴴ, 1, 1) @@ -265,7 +270,12 @@ function svd_trunc_pullback!( else Y = rmul!(AP' * X₀, Diagonal(S⁻¹)) Y .+= Y₀ᴴ' - Y = accelerative_smith_iteration!(Y, similar(Y), AP' * AP, S⁻¹ .^ 2, degeneracy_atol, maxiter) + APZ = similar(Y, m, p) # for applying AP' AP without forming it + Y = hermitian_stein_cg!( + Y, (GZ, Z) -> mul!(GZ, AP', mul!(view(APZ, :, axes(Z, 2)), AP, Z)), () -> AP' * AP, + S⁻¹ .^ 2, degeneracy_atol, maxiter; + cost_apply = 2 * m * n, cost_apply_formed = n^2, cost_form = n^2 * m + ) X = rmul!(AP * Y, Diagonal(S⁻¹)) X .+= X₀ ΔA = mul!(ΔA, X, Vᴴ, 1, 1)