Skip to content

[WIP] Solve the truncated svd and eigh pullbacks by conjugate gradients - #296

Draft
leburgel wants to merge 1 commit into
mainfrom
lb/trunc_pullback_cg
Draft

leburgel wants to merge 1 commit into
mainfrom
lb/trunc_pullback_cg

Conversation

@leburgel

@leburgel leburgel commented Oct 4, 2026

Copy link
Copy Markdown
Member

Draft.

For every kept column i, svd_trunc_pullback! and eigh_trunc_pullback! solve (1 - wᵢ G) xᵢ = bᵢ with G Hermitian (SVD: G = P Pᴴ or Pᴴ P, whichever is smaller, and wᵢ = 1/σᵢ²; eigh: G = P and wᵢ = 1/λᵢ). main solves it by doubling, at O(n³) per squaring of G. This PR uses conjugate gradients on all columns at once, at O(n² k) per iteration for k kept columns, which pays off when k ≪ n (which is common in the context of CTMRG runs in PEPSKit.jl, which motivated this investigation). There is no new keyword.

  • The conjugate-gradient loop follows KrylovKit's CG, batched over the columns.
  • A column stops once its residual is below degeneracy_atol times the smallest Ritz value, so that its error is about degeneracy_atol, as for doubling.
  • maxiter now caps the conjugate-gradient iterations (default 10 times the dimension), with a warning if they do not converge.
  • SVD: G is applied as two products until forming it pays off. eigh: G = P is indefinite, so long runs restart on (1 - wᵢ² P²) xᵢ = (1 + wᵢ P) bᵢ, which needs about half the iterations.
  • The solver, hermitian_stein_cg!, goes into a new file src/common/stein.jl, together with accelerative_smith_iteration!, moved there unchanged (still used by eig_trunc_pullback!).

Timed against main's solver on the same equations, in one process with the calls interleaved, on 204 cases in real and complex arithmetic (script below, BENCH_MINTOTAL=1):

set dec speed-up: median (range) slower largest error (PR / main)
CTMRG-like svd 3.45× (1.32–7.35) 0/8 1.0e-14 / 1.0e-14
CTMRG-like eigh 5.71× (1.53–9.45) 0/8 2.5e-13 / 2.5e-13
sweep svd 2.25× (0.97–8.00) 1/64 7.7e-12 / 7.8e-12
sweep eigh 2.36× (1.02–9.82) 0/64 2.3e-09 / 2.3e-09
randn svd 1.19× (0.68–3.04) 9/30 3.3e-13 / 6.5e-13
randn eigh 1.25× (0.64–2.71) 9/30 7.1e-13 / 1.1e-12
  • CTMRG-like: n = χD² for (D, χ) from (3, 20) to (6, 80), k = χ, values decaying by 0.98 (SVD) or 0.956 (eigh) per value. Sweep: n = 400, 1200, k/n from 1/64 to 1/2, ratio at the cut from 0.9 to 0.9999. randn: n = 100, 400, 1000, k from n/19 to 17n/19.
  • Most of the slower cases are n = 100 (calls of about 1 ms, 0.64–0.83×); for n ≥ 400 the slowest is 0.72×. Per-case speed-ups vary by up to about 2× between runs; the set medians agreed within 11%.

eig_trunc_pullback! keeps doubling. Its G = APᴴ is not Hermitian, so conjugate gradients do not apply, and for eigenvalues spread over a disk no Krylov method converges faster than the Neumann series itself.

Benchmark script

Run with this branch developed; the numbers above used BENCH_MINTOTAL=1 (28 minutes).

# Time svd_trunc_pullback! / eigh_trunc_pullback! of this PR ("auto") against main's solver, forced
# on this branch ("dbl": `accelerative_smith_iteration!` on the full right-hand side), in one process with the
# calls interleaved, on CTMRG-like spectra, a sweep of the ratio at the cut and random matrices.
# Error relative to the full pullback with the truncation indices. BENCH_MINTOTAL: seconds per solver.
using LinearAlgebra, Random, Printf
using MatrixAlgebraKit
using MatrixAlgebraKit: diagview, findtruncated, svd_pullback!, svd_trunc_pullback!,
    eigh_pullback!, eigh_trunc_pullback!, remove_svd_gauge_dependence!, remove_eigh_gauge_dependence!
BLAS.set_num_threads(parse(Int, get(ENV, "BENCH_BLAS", "4")))
const SETS = split(get(ENV, "BENCH_SETS", "ctmrg,sweep,scaled"), ",")
const MINTOTAL = parse(Float64, get(ENV, "BENCH_MINTOTAL", "3"))
const MODE = Ref(:auto)
# "dbl": main's solver, `accelerative_smith_iteration!` (maxiter 100, main's default), on the same equation
@eval MatrixAlgebraKit function hermitian_stein_cg!(X::AbstractMatrix, applyG!, formG, w::AbstractVector, atol::Real, maxiter::Int; kwargs...)
    Main.MODE[] === :dbl && return accelerative_smith_iteration!(X, similar(X), formG(), copy(w), atol, 100)
    return invoke(hermitian_stein_cg!, Tuple{Any, Any, Any, Any, Any, Any}, X, applyG!, formG, w, atol, maxiter; kwargs...)
end
const MODES = (:auto, :dbl)
# cases to run, by position ("all" by default)
const CASES = get(ENV, "BENCH_CASES", "all") == "all" ? Set(1:10_000) : Set(parse.(Int, split(ENV["BENCH_CASES"], ",")))
const COUNT = Ref(0)
function interleaved(f)
    T = Dict(m => Float64[] for m in MODES)
    for m in keys(T)
        MODE[] = m; f()
    end
    for _ in 1:200
        for m in MODES
            MODE[] = m; s = time_ns(); f(); push!(T[m], (time_ns() - s) / 1.0e9)
        end
        all(sum(t) > MINTOTAL for t in values(T)) && break
    end
    MODE[] = :auto
    return T
end
median(x) = (y = sort(x); n = length(y); isodd(n) ? y[(n + 1) ÷ 2] : (y[n ÷ 2] + y[n ÷ 2 + 1]) / 2)
relerr(a, b) = norm(a - b) / norm(b)
function out(set, dec, A, p, ρ, f, ref)
    T = interleaved(f)
    e = Dict{Symbol, Float64}()
    for m in MODES
        MODE[] = m; e[m] = relerr(f(), ref)
    end
    MODE[] = :auto
    @printf("RETIME case=%d set=%s dec=%s T=%s n=%d p=%d rho=%.4f rounds=%d", COUNT[], set, dec, eltype(A), size(A, 1), p, ρ, length(T[:auto]))
    for m in MODES
        @printf(" %s_min=%.3e %s_med=%.3e err_%s=%.2e", m, minimum(T[m]), m, median(T[m]), m, e[m])
    end
    println()
    flush(stdout)
end
function run_svd(set, A, ind, rng)
    if !((COUNT[] += 1) in CASES)
        randn!(rng, similar(A)); randn!(rng, similar(A)); randn(rng, real(eltype(A)), size(A, 1))
        return
    end
    U, S, Vᴴ = svd_compact(A)
    ΔU, ΔVᴴ = remove_svd_gauge_dependence!(randn!(rng, similar(U)), randn!(rng, similar(Vᴴ)), U, S, Vᴴ)
    ΔS = Diagonal(randn(rng, real(eltype(A)), size(S, 1)))
    trunc = (U[:, ind], Diagonal(diagview(S)[ind]), Vᴴ[ind, :])
    Δt = (ΔU[:, ind], Diagonal(diagview(ΔS)[ind]), ΔVᴴ[ind, :])
    p = length(ind); ρ = diagview(S)[p + 1] / diagview(S)[p]
    ref = svd_pullback!(zero(A), A, (U, S, Vᴴ), Δt, ind)
    f() = svd_trunc_pullback!(zero(A), A, trunc, Δt)
    out(set, "svd", A, p, ρ, f, ref)
end
function run_eigh(set, H, ind, rng)
    if !((COUNT[] += 1) in CASES)
        randn!(rng, similar(H)); randn(rng, real(eltype(H)), size(H, 1))
        return
    end
    D, V = eigh_full(H)
    ΔV = remove_eigh_gauge_dependence!(randn!(rng, similar(V)), D, V)
    ΔD = Diagonal(randn(rng, real(eltype(H)), size(D, 1)))
    trunc = (Diagonal(diagview(D)[ind]), V[:, ind]); Δt = (Diagonal(diagview(ΔD)[ind]), ΔV[:, ind])
    rest = setdiff(axes(D, 1), ind); p = length(ind)
    ρ = maximum(abs, diagview(D)[rest]) / minimum(abs, diagview(D)[ind])
    ref = eigh_pullback!(zero(H), H, (D, V), Δt, ind)
    f() = eigh_trunc_pullback!(zero(H), H, trunc, Δt)
    out(set, "eigh", H, p, ρ, f, ref)
end
println("MatrixAlgebraKit from ", pathof(MatrixAlgebraKit), ", BLAS threads ", BLAS.get_num_threads())
rng = Xoshiro(123)   # matrices
crng = Xoshiro(456)  # cotangents
for T in (Float64, ComplexF64)
    if "ctmrg" in SETS
        for (D, χ) in ((3, 20), (4, 40), (5, 60), (6, 80))
            n = χ * D^2
            Q1, Q2 = Matrix(qr(randn(rng, T, n, n)).Q), Matrix(qr(randn(rng, T, n, n)).Q)
            A = Q1 * Diagonal(max.(0.98 .^ (0:(n - 1)), 1.0e-10)) * Q2'
            run_svd("ctmrg", A, findtruncated(svdvals(A), truncrank(χ)), crng)
            λ = max.(0.956 .^ (0:(n - 1)), 1.0e-10) .* rand(rng, (-1, 1), n)
            H = Matrix(Hermitian(Q1 * Diagonal(λ) * Q1'))
            run_eigh("ctmrg", H, findtruncated(eigh_vals(H), truncrank(χ; by = abs)), crng)
        end
    end
    if "sweep" in SETS
        for n in (400, 1200), f in (64, 16, 4, 2), ρc in (0.9, 0.99, 0.999, 0.9999)
            p = n ÷ f
            σ = vcat(0.99 .^ (0:(p - 1)), 0.99^(p - 1) * ρc .* 0.99 .^ (0:(n - p - 1)))
            Q1, Q2 = Matrix(qr(randn(rng, T, n, n)).Q), Matrix(qr(randn(rng, T, n, n)).Q)
            run_svd("sweep", Q1 * Diagonal(σ) * Q2', collect(1:p), crng)
            H = Matrix(Hermitian(Q1 * Diagonal(σ .* rand(rng, (-1, 1), n)) * Q1'))
            run_eigh("sweep", H, findtruncated(eigh_vals(H), truncrank(p; by = abs)), crng)
        end
    end
    if "scaled" in SETS
        for m in (100, 400, 1000)
            A = randn(rng, T, m, m)
            H = Matrix(Hermitian(randn(rng, T, m, m)))
            for f in (1, 5, 9, 13, 17)
                r = max(1, round(Int, f / 19 * m))
                run_svd("scaled", A, findtruncated(svdvals(A), truncrank(r)), crng)
                run_eigh("scaled", H, findtruncated(eigh_vals(H), truncrank(r; by = abs)), crng)
            end
        end
    end
end

@leburgel
leburgel marked this pull request as draft October 4, 2026 08:08
@leburgel
leburgel force-pushed the lb/trunc_pullback_cg branch from 1a38822 to ffa3c6a Compare October 4, 2026 08:19
@codecov

codecov Bot commented Oct 4, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 97.45763% with 3 lines in your changes missing coverage. Please review.

Files with missing lines Patch % Lines
src/common/stein.jl 97.08% 3 Missing ⚠️
Files with missing lines Coverage Δ
src/MatrixAlgebraKit.jl 100.00% <ø> (ø)
src/common/pullbacks.jl 100.00% <ø> (+7.14%) ⬆️
src/pullbacks/eigh.jl 86.25% <100.00%> (+0.17%) ⬆️
src/pullbacks/svd.jl 94.59% <100.00%> (+0.11%) ⬆️
src/common/stein.jl 97.08% <97.08%> (ø)
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant