Skip to content

Make svd_pullback! and eigh_pullback! cost proportional to the number of cotangent columns - #292

Open
leburgel wants to merge 3 commits into
mainfrom
lb/full_pullback_kept_columns
Open

leburgel wants to merge 3 commits into
mainfrom
lb/full_pullback_kept_columns

Conversation

@leburgel

@leburgel leburgel commented Oct 1, 2026

Copy link
Copy Markdown
Member

When the cotangents of svd_pullback! are given on only k of the r singular vectors (through ind), the pullback pads them with zeros to all r columns and continues with r × r matrices: check_and_prepare_svd_cotangents forms U₁' * ΔU₁ and ΔU₁ - U₁ * (U₁' * ΔU₁), the same for V, and the result is applied as U₁ * (UᴴΔAV * V₁ᴴ). The cost is therefore O(m n r) whatever k is. eigh_pullback! does the same with V' * ΔV₁ and V * VᴴΔAV * V', at cost O(n³) for an n × n matrix. This is the common case for a full pullback of a truncated decomposition, where ind holds the kept indices.

With cotangents on the columns K only, U₁ᴴΔU₁ and V₁ᴴΔV₁ are nonzero only in the columns K. Then UᴴΔAV is nonzero only in the rows and columns K, and its rows follow from its columns by antihermiticity. In this PR, check_and_prepare_svd_cotangents therefore no longer pads the cotangents, but computes only the r × k block of columns K of UᴴΔAV and the corresponding block of its rows. For 2k ≤ r, svd_pullback! applies these two blocks directly as rank-k updates, at cost O(m n k); otherwise it assembles UᴴΔAV from them and applies it as before. check_and_prepare_eigh_cotangents and eigh_pullback! are changed in the same way, so that for 2k ≤ n the cost of eigh_pullback! is O(n² k) instead of O(n³). The gauge check covers the same entries as before, since the columns K contain every nonzero entry up to conjugation. Cotangents on columns beyond a full rank have components along all of U₁ or V₁ᴴ, so in that case the block still spans all r columns. svd_trunc_pullback! and eigh_trunc_pullback! call the same functions with all columns and receive the same full matrices as before.

Bug fix. This PR also fixes svd_pullback! for a nonzero ΔS with an ind other than 1:k. Since #232, check_and_prepare_svd_cotangents on main indexes ΔS by column number instead of by position within ind, so that for example ind = [3, 1, 7, 2] throws a BoundsError. The new code adds each entry of ΔS to the diagonal entry of its own column; the comparison below includes this case.

On random real and complex, square and rectangular, full-rank and rank-deficient matrices, with ind = 1:5, [3, 1, 7, 2] and 1:6, and with zero ΔU, ΔVᴴ or ΔS, the result agrees with main (called with the same cotangents zero-padded to all columns) to 7.9e-16 relative. On square matrices with exponentially decaying singular values or eigenvalues and k between n / 36 and n / 9 (the spectra and sizes of a CTMRG step in PEPSKit.jl, with n = χD² and k = χ, which motivated this investigation), it is 4–22× faster:

decomposition eltype n k main (s) this PR (s) speedup
SVD Float64 180 20 1.30e-3 3.29e-4 3.9×
SVD Float64 640 40 4.57e-2 5.51e-3 8.3×
SVD Float64 1500 60 0.515 2.34e-2 22.0×
SVD Float64 2880 80 2.31 0.173 13.3×
SVD ComplexF64 180 20 4.40e-3 8.23e-4 5.3×
SVD ComplexF64 640 40 0.175 1.67e-2 10.5×
SVD ComplexF64 1500 60 1.10 7.34e-2 14.9×
SVD ComplexF64 2880 80 9.12 0.622 14.7×
eigh Float64 180 20 8.06e-4 1.97e-4 4.1×
eigh Float64 640 40 2.30e-2 3.23e-3 7.1×
eigh Float64 1500 60 0.276 2.34e-2 11.8×
eigh Float64 2880 80 1.58 0.121 13.1×
eigh ComplexF64 180 20 2.27e-3 5.30e-4 4.3×
eigh ComplexF64 640 40 8.68e-2 1.04e-2 8.3×
eigh ComplexF64 1500 60 0.894 8.92e-2 10.0×
eigh ComplexF64 2880 80 4.72 0.240 19.6×

Minimum times on a laptop with 4 BLAS threads. main's src/pullbacks/svd.jl and src/pullbacks/eigh.jl were loaded into the same process, and the two versions were timed alternately with the same arguments.

On random n × n matrices (n = 400 and 1000, real and complex; same timing method), the speedup over main for cotangents on the first k columns (for eigh, the k eigenvalues of largest magnitude) is:

k / n path SVD eigh
1/20 new 9.3–14× 7.5–16×
1/4 new 2.7–3.3× 2.3–2.4×
9/20 new 1.55–1.67× 1.30–1.76×
1/2 new 1.37–1.50× 1.19–1.23×
11/20 dense 1.34–1.53× 1.15–1.18×
7/10 dense 0.76–1.21× 0.94–1.25×
9/10 dense 0.77–1.10× 0.97–1.04×
1 (ind = Colon()) dense 0.99–1.14× 1.00–1.16×

With 2k ≤ r the new path is faster in every case. Above that, the r × r matrix is formed from the k computed columns and applied as before, which is slower than main in five of the 24 cases with 2k > r: by up to 1.3× in two SVD cases at n = 1000, and by at most 6% in the other three. With ind = Colon() the result is identical to that of main.

Benchmark
using LinearAlgebra, Random, Printf
using MatrixAlgebraKit
using MatrixAlgebraKit: svd_pullback!, eigh_pullback!, diagview

# gauge-invariant cotangents on the columns `ind` (the imaginary part of diag(Xᴴ ΔX) is the gauge part)
noimagdiag!(ΔX, X) = (ΔX .-= X .* transpose(im .* imag.(diag(X' * ΔX))); ΔX)
noimagdiag!(ΔX::AbstractMatrix{<:Real}, X) = ΔX
pad(X, ind, p) = (Y = zeros(eltype(X), size(X, 1), p); Y[:, ind] .= X; Y)
padvec(x, ind, p) = (y = zeros(eltype(x), p); y[ind] .= x; y)

# time of the pullback with cotangents on the columns ind, and its difference with the same
# cotangents zero-padded to all columns
function timeit(name, pullback!, A, X, cots, padded, ind)
    f = () -> pullback!(zero(A), A, X, cots, ind)
    t = (f(); minimum(@elapsed(f()) for _ in 1:5))
    ref = pullback!(zero(A), A, X, padded)
    @printf("%-4s %-10s n = %4d  k = %2d  time %.2e s  difference %.1e\n",
        name, eltype(A), size(A, 1), length(ind), t, norm(f() - ref) / norm(ref))
end

BLAS.set_num_threads(4)
rng = Xoshiro(1)
for T in (Float64, ComplexF64), (n, k) in ((180, 20), (640, 40), (1500, 60), (2880, 80))
    ind = 1:k
    Q₁, Q₂ = Matrix(qr(randn(rng, T, n, n)).Q), Matrix(qr(randn(rng, T, n, n)).Q)

    A = Q₁ * Diagonal(max.(0.98 .^ (0:(n - 1)), 1.0e-10)) * Q₂'
    U, S, Vᴴ = svd_compact(A)
    ΔU = noimagdiag!(randn(rng, T, n, k), U[:, ind])
    ΔV = noimagdiag!(randn(rng, T, n, k), Vᴴ[ind, :]')
    ΔS = randn(rng, real(T), k)
    timeit("svd", svd_pullback!, A, (U, S, Vᴴ), (ΔU, Diagonal(ΔS), copy(ΔV')),
        (pad(ΔU, ind, n), Diagonal(padvec(ΔS, ind, n)), copy(pad(ΔV, ind, n)')), ind)

    λ = max.(0.956 .^ (0:(n - 1)), 1.0e-10) .* rand(rng, (-1, 1), n)
    H = Matrix(Hermitian(Q₁ * Diagonal(λ) * Q₁'))
    D, V = eigh_full(H)
    indh = sort(sortperm(abs.(diagview(D)); rev = true)[1:k])
    ΔV = noimagdiag!(randn(rng, T, n, k), V[:, indh])
    ΔD = randn(rng, real(T), k)
    timeit("eigh", eigh_pullback!, H, (D, V), (Diagonal(ΔD), ΔV),
        (Diagonal(padvec(ΔD, indh, n)), pad(ΔV, indh, n)), indh)
end

On main:

svd  Float64    n =  180  k = 20  time 1.96e-03 s  difference 0.0e+00
eigh Float64    n =  180  k = 20  time 5.86e-04 s  difference 0.0e+00
svd  Float64    n =  640  k = 40  time 2.66e-02 s  difference 0.0e+00
eigh Float64    n =  640  k = 40  time 1.22e-02 s  difference 0.0e+00
svd  Float64    n = 1500  k = 60  time 3.78e-01 s  difference 0.0e+00
eigh Float64    n = 1500  k = 60  time 2.88e-01 s  difference 0.0e+00
svd  Float64    n = 2880  k = 80  time 3.08e+00 s  difference 0.0e+00
eigh Float64    n = 2880  k = 80  time 1.27e+00 s  difference 0.0e+00
svd  ComplexF64 n =  180  k = 20  time 2.18e-03 s  difference 0.0e+00
eigh ComplexF64 n =  180  k = 20  time 2.29e-03 s  difference 0.0e+00
svd  ComplexF64 n =  640  k = 40  time 1.79e-01 s  difference 0.0e+00
eigh ComplexF64 n =  640  k = 40  time 8.61e-02 s  difference 0.0e+00
svd  ComplexF64 n = 1500  k = 60  time 1.83e+00 s  difference 0.0e+00
eigh ComplexF64 n = 1500  k = 60  time 5.79e-01 s  difference 0.0e+00
svd  ComplexF64 n = 2880  k = 80  time 9.11e+00 s  difference 0.0e+00
eigh ComplexF64 n = 2880  k = 80  time 3.91e+00 s  difference 0.0e+00

With this PR:

svd  Float64    n =  180  k = 20  time 3.24e-04 s  difference 1.1e-15
eigh Float64    n =  180  k = 20  time 1.98e-04 s  difference 7.4e-16
svd  Float64    n =  640  k = 40  time 5.45e-03 s  difference 1.2e-15
eigh Float64    n =  640  k = 40  time 3.45e-03 s  difference 7.2e-16
svd  Float64    n = 1500  k = 60  time 4.00e-02 s  difference 1.1e-15
eigh Float64    n = 1500  k = 60  time 2.41e-02 s  difference 7.1e-16
svd  Float64    n = 2880  k = 80  time 2.17e-01 s  difference 1.0e-15
eigh Float64    n = 2880  k = 80  time 8.40e-02 s  difference 7.7e-16
svd  ComplexF64 n =  180  k = 20  time 7.32e-04 s  difference 1.4e-15
eigh ComplexF64 n =  180  k = 20  time 5.02e-04 s  difference 7.8e-16
svd  ComplexF64 n =  640  k = 40  time 1.39e-02 s  difference 1.5e-15
eigh ComplexF64 n =  640  k = 40  time 4.62e-03 s  difference 9.3e-16
svd  ComplexF64 n = 1500  k = 60  time 1.39e-01 s  difference 1.4e-15
eigh ComplexF64 n = 1500  k = 60  time 4.81e-02 s  difference 9.9e-16
svd  ComplexF64 n = 2880  k = 80  time 6.47e-01 s  difference 1.2e-15
eigh ComplexF64 n = 2880  k = 80  time 2.13e-01 s  difference 9.4e-16

@leburgel
leburgel requested a review from Jutho October 1, 2026 11:53
Comment thread src/common/pullbacks.jl
Comment on lines +39 to +52
"""
antihermitian_columns!(X, K)

Given the columns `K` of a square matrix that is nonzero only in its columns `K`, overwrite `X`
with the same columns of the antihermitian part of that matrix.
"""
function antihermitian_columns!(X, K)
# NOTE: all columns in order (e.g. `ind = Colon()`): the original in-place projection
is_leading_index(K, size(X, 1)) && return project_antihermitian!(X)
XKK = project_antihermitian!(X[K, :])
X ./= 2
X[K, :] .= XKK
return X
end

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.

Any reason to use K instead of ind as variable name? I find this quite confusing as I often denote certain matrices coming up in the pullbacks as K (i.e. the antihermitian matrix associated with the infinitesimal in-space rotation of an isometry).

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.

Also, I don't understand what is happening here. Is X square, or is X already the K columns of a larger square matrix. And then X[K, :] is square, corresponding to the diagonal block of that matrix? I think it is the latter, but the doc string could be a bit clearer.

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.

No good reason at all, switched it back to use ind. I also updated the docstring and variable names to make it more clear what's actually happening. X holds the columns ind of a larger square matrix M that is nonzero only in those columns, and X[ind, :] = M[ind, ind] is its diagonal block, the only part of X to which M' contributes.

@codecov

codecov Bot commented Oct 1, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.

Files with missing lines Coverage Δ
src/common/pullbacks.jl 100.00% <100.00%> (ø)
src/pullbacks/eigh.jl 87.27% <100.00%> (+1.13%) ⬆️
src/pullbacks/svd.jl 94.25% <100.00%> (+0.89%) ⬆️
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

Comment thread src/pullbacks/eigh.jl Outdated
indD = select_indices(axes(D, 1), ind)
indV = select_indices(axes(V, 2), ind)
K = select_indices(axes(D, 1), ind)
k = length(K)

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.

Can we keep the ind or indD name?

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.

Or ind′ in case just want a single variable, and since ind is already in use.

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.

I renamed to ind′ for consistency.

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.

2 participants