Skip to content
Draft
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
1 change: 1 addition & 0 deletions src/MatrixAlgebraKit.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
36 changes: 0 additions & 36 deletions src/common/pullbacks.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
177 changes: 177 additions & 0 deletions src/common/stein.jl
Original file line number Diff line number Diff line change
@@ -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`.
"""

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is far from a "simple CG". It will take me quite some time to review it. I also don't particularly like the coding style, though I don't know to pinpoint that more precisely, or thus, what action could improve it 😄 . Maybe after I understand the code I will have some suggestions.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This was more a placeholder for pointing out that CG is a good Stein solver for these specific cases. Other than requesting it to be "as close to the KrylovKit CG as possible", I didn't really check the actual implementation. I'm perfectly fine with replacing the entire solver, I just didn't have the expertise to write a good one myself.

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
9 changes: 6 additions & 3 deletions src/pullbacks/eigh.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down
16 changes: 13 additions & 3 deletions src/pullbacks/svd.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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ᴴ
Expand Down Expand Up @@ -257,15 +257,25 @@ 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)
Δ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)
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)
Expand Down
Loading