Skip to content

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

Merged
leburgel merged 24 commits into
mainfrom
lb/full_pullback_kept_columns
Oct 8, 2026
Merged

leburgel merged 24 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 p 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 p 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 ind only, U₁ᴴΔU₁ and V₁ᴴΔV₁ are nonzero only in the columns ind. Then UᴴΔAV is nonzero only in the rows and columns ind, 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 × p block of columns ind of UᴴΔAV and the corresponding block of its rows. For 2p ≤ r, svd_pullback! applies these two blocks directly as rank-p updates, at cost O(m n p); 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² p) instead of O(n³). The gauge check covers the same entries as before, since the columns ind 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:p. 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 p between n / 36 and n / 9 (the spectra and sizes of a CTMRG step in PEPSKit.jl, with n = χD² and p = χ, which motivated this investigation), it is 4–22× faster:

decomposition eltype n p 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 p columns (for eigh, the p eigenvalues of largest magnitude) is:

p / 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 2p ≤ r the new path is faster in every case. Above that, the r × r matrix is formed from the p computed columns and applied as before, which is slower than main in five of the 24 cases with 2p > 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  p = %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, p) in ((180, 20), (640, 40), (1500, 60), (2880, 80))
    ind = 1:p
    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, p), U[:, ind])
    ΔV = noimagdiag!(randn(rng, T, n, p), Vᴴ[ind, :]')
    ΔS = randn(rng, real(T), p)
    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:p])
    ΔV = noimagdiag!(randn(rng, T, n, p), V[:, indh])
    ΔD = randn(rng, real(T), p)
    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  p = 20  time 1.96e-03 s  difference 0.0e+00
eigh Float64    n =  180  p = 20  time 5.86e-04 s  difference 0.0e+00
svd  Float64    n =  640 p = 40  time 2.66e-02 s  difference 0.0e+00
eigh Float64    n =  640  p = 40  time 1.22e-02 s  difference 0.0e+00
svd  Float64    n = 1500  p = 60  time 3.78e-01 s  difference 0.0e+00
eigh Float64    n = 1500  p = 60  time 2.88e-01 s  difference 0.0e+00
svd  Float64    n = 2880  p = 80  time 3.08e+00 s  difference 0.0e+00
eigh Float64    n = 2880  p = 80  time 1.27e+00 s  difference 0.0e+00
svd  ComplexF64 n =  180  p = 20  time 2.18e-03 s  difference 0.0e+00
eigh ComplexF64 n =  180  p = 20  time 2.29e-03 s  difference 0.0e+00
svd  ComplexF64 n =  640  p = 40  time 1.79e-01 s  difference 0.0e+00
eigh ComplexF64 n =  640  p = 40  time 8.61e-02 s  difference 0.0e+00
svd  ComplexF64 n = 1500  p = 60  time 1.83e+00 s  difference 0.0e+00
eigh ComplexF64 n = 1500  p = 60  time 5.79e-01 s  difference 0.0e+00
svd  ComplexF64 n = 2880  p = 80  time 9.11e+00 s  difference 0.0e+00
eigh ComplexF64 n = 2880  p = 80  time 3.91e+00 s  difference 0.0e+00

With this PR:

svd  Float64    n =  180  p = 20  time 3.24e-04 s  difference 1.1e-15
eigh Float64    n =  180 p = 20  time 1.98e-04 s  difference 7.4e-16
svd  Float64    n =  640  p = 40  time 5.45e-03 s  difference 1.2e-15
eigh Float64    n =  640  p = 40  time 3.45e-03 s  difference 7.2e-16
svd  Float64    n = 1500  p = 60  time 4.00e-02 s  difference 1.1e-15
eigh Float64    n = 1500  p = 60  time 2.41e-02 s  difference 7.1e-16
svd  Float64    n = 2880  p = 80  time 2.17e-01 s  difference 1.0e-15
eigh Float64    n = 2880  p = 80  time 8.40e-02 s  difference 7.7e-16
svd  ComplexF64 n =  180  p = 20  time 7.32e-04 s  difference 1.4e-15
eigh ComplexF64 n =  180  p = 20  time 5.02e-04 s  difference 7.8e-16
svd  ComplexF64 n =  640  p = 40  time 1.39e-02 s  difference 1.5e-15
eigh ComplexF64 n =  640  p = 40  time 4.62e-03 s  difference 9.3e-16
svd  ComplexF64 n = 1500  p = 60  time 1.39e-01 s  difference 1.4e-15
eigh ComplexF64 n = 1500  p = 60  time 4.81e-02 s  difference 9.9e-16
svd  ComplexF64 n = 2880  p = 80  time 6.47e-01 s  difference 1.2e-15
eigh ComplexF64 n = 2880  p = 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
@codecov

codecov Bot commented Oct 1, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 98.26087% with 2 lines in your changes missing coverage. Please review.

Files with missing lines Patch % Lines
src/pullbacks/svd.jl 97.22% 2 Missing ⚠️
Files with missing lines Coverage Δ
src/common/pullbacks.jl 94.11% <100.00%> (+1.26%) ⬆️
src/pullbacks/eigh.jl 86.90% <100.00%> (+2.09%) ⬆️
src/pullbacks/lq.jl 98.98% <100.00%> (+2.02%) ⬆️
src/pullbacks/qr.jl 98.98% <100.00%> (+2.02%) ⬆️
src/pullbacks/svd.jl 95.02% <97.22%> (+19.88%) ⬆️

... and 5 files with indirect coverage changes

🚀 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
Comment thread src/pullbacks/eigh.jl Outdated
Comment thread src/pullbacks/eigh.jl Outdated
Comment thread src/pullbacks/eigh.jl Outdated
@leburgel
leburgel force-pushed the lb/full_pullback_kept_columns branch from 80d4af2 to b0a90bd Compare October 2, 2026 14:21
Comment thread src/pullbacks/eigh.jl Outdated
Comment thread src/pullbacks/svd.jl Outdated
@leburgel
leburgel force-pushed the lb/full_pullback_kept_columns branch from b0a90bd to fc18c56 Compare October 3, 2026 07:32
@leburgel

leburgel commented Oct 3, 2026 •

Copy link
Copy Markdown
Member Author

The views of U, Vᴴ and V broke the GPU tests (the failing Buildkite mooncake jobs): when ind′ is a Vector{Int}, the view is not a StridedMatrix, so mul! falls back to the generic matrix product, which is scalar indexing on a CuArray. I therefore went back to copies, renamed to Uₖ, Vᴴₖ and Vₖ like Dₖ, and kept the view of S, which is only broadcast.

They were not the best idea on CPU either, because a Vector ind′ is common: truncerror, and truncrank for eigh, return one even when the kept columns are contiguous. The ratio of times views/copies for the pullbacks (laptop, 4 BLAS threads, Float64 and ComplexF64, svd of 1200×1000, eigh of 1000×1000, p = 50, 200, 700):

ind a range ind a Vector
svd_pullback! 0.98–1.18 1.7–7.0
eigh_pullback! 0.99–1.01 3.3–9.5 (1.0 for p > n/2, which does not use Vₖ)

The copies only cost O(n p) next to the O(n² p) products.

Benchmark script (run on fc18c56, the version with views)
# Views vs copies of the kept columns in svd_pullback!/eigh_pullback! (untracked, for the review reply):
#   1. mul! with a view of the columns ind as one factor, for ind a range and a Vector
#   2. the pullbacks of this branch (views of U, Vᴴ and V) against the same code with copies
#      `U[:, ind′]`, `Vᴴ[ind′, :]`, `V[:, ind′]` (module Copies), timed alternately in one process
using LinearAlgebra, Random, Printf
using MatrixAlgebraKit
using MatrixAlgebraKit: svd_pullback!, eigh_pullback!
BLAS.set_num_threads(4)
module Copies
using LinearAlgebra, MatrixAlgebraKit
const SUBS = ("view(U, :, ind′)" => "U[:, ind′]", "view(Vᴴ, ind′, :)" => "Vᴴ[ind′, :]", "view(V, :, ind′)" => "V[:, ind′]")
const SRC = map(("svd", "eigh")) do f
    src = read(joinpath(pkgdir(MatrixAlgebraKit), "src", "pullbacks", "$f.jl"), String)
    any(occursin(first(s), src) for s in SUBS) || error("no view in $f.jl")
    return replace(src, SUBS...)
end
const OWN = Set(Symbol(something(m[1], m[2])) for src in SRC for m in eachmatch(r"^\s*function ([^\s(]+)\(|^(?:function )?([A-Za-z_][^\s(]*)\([^\n]*\)\s*=[^=]"m, src))
for name in names(MatrixAlgebraKit; all = true)
    (name in OWN || startswith(string(name), "#") || isdefined(@__MODULE__, name)) && continue
    @eval const $name = MatrixAlgebraKit.$name
end
foreach(src -> include_string(@__MODULE__, src), SRC)
end
function mintimes(f, g; mintotal = 2.0, minreps = 5, maxreps = 15)
    f(); g(); tf = tg = Inf; tot = 0.0
    for i in 1:maxreps
        s = time_ns(); f(); df = (time_ns() - s) / 1.0e9
        s = time_ns(); g(); dg = (time_ns() - s) / 1.0e9
        tf, tg = min(tf, df), min(tg, dg); tot += df + dg
        i >= minreps && tot > mintotal && break
    end
    return tf, tg
end
noimagdiag!(ΔX, X) = (ΔX .-= X .* transpose(im .* imag.(diag(X' * ΔX))); ΔX)
noimagdiag!(ΔX::AbstractMatrix{<:Real}, X) = ΔX
indkinds(n, p, rng) = (("range", 1:p), ("vector", collect(1:p)), ("perm", randperm(rng, n)[1:p]))

println("== mul!(C, A, B') with B = view(V, :, ind) or V[:, ind], n = 1000, BLAS threads 4")
let rng = Xoshiro(1), n = 1000
    for T in (Float64, ComplexF64), p in (50, 200)
        A = randn(rng, T, n, p); V = randn(rng, T, n, n); C = zeros(T, n, n)
        for (name, ind) in indkinds(n, p, rng)
            tv, tc = mintimes(() -> mul!(C, A, view(V, :, ind)', 1, 1), () -> mul!(C, A, V[:, ind]', 1, 1))
            @printf("MUL T=%s p=%d ind=%s view=%.3e copy=%.3e view/copy=%.1f\n", T, p, name, tv, tc, tv / tc)
        end
    end
end

println("== pullbacks: views (branch) vs copies, svd m=1200 n=1000, eigh n=1000")
let rng = Xoshiro(2), m = 1200, n = 1000
    for T in (Float64, ComplexF64)
        A = randn(rng, T, m, n); U, S, Vᴴ = svd_compact(A); Sd = diag(S)
        H = randn(rng, T, n, n); H = H + H'; D, V = eigh_full(H); Dd = diag(D)
        for p in (50, 200, 700), (name, ind) in indkinds(n, p, rng)
            ΔU = noimagdiag!(randn(rng, T, m, p), U[:, ind])
            ΔVᴴ = copy(noimagdiag!(randn(rng, T, n, p), Vᴴ[ind, :]')')
            ΔS = Diagonal(randn(rng, real(T), p))
            f(pb) = () -> pb(zero(A), A, (U, S, Vᴴ), (ΔU, ΔS, ΔVᴴ), ind)
            err = norm(f(svd_pullback!)() - f(Copies.svd_pullback!)())
            tv, tc = mintimes(f(svd_pullback!), f(Copies.svd_pullback!))
            @printf("PB dec=svd T=%s p=%d ind=%s views=%.3e copies=%.3e views/copies=%.2f diff=%.0e\n", T, p, name, tv, tc, tv / tc, err)
            ΔV = noimagdiag!(randn(rng, T, n, p), V[:, ind]); ΔD = Diagonal(randn(rng, real(T), p))
            g(pb) = () -> pb(zero(H), H, (D, V), (ΔD, ΔV), ind)
            err = norm(g(eigh_pullback!)() - g(Copies.eigh_pullback!)())
            tv, tc = mintimes(g(eigh_pullback!), g(Copies.eigh_pullback!))
            @printf("PB dec=eigh T=%s p=%d ind=%s views=%.3e copies=%.3e views/copies=%.2f diff=%.0e\n", T, p, name, tv, tc, tv / tc, err)
        end
    end
end

@leburgel
leburgel force-pushed the lb/full_pullback_kept_columns branch from 64e4a1d to bad2b2c Compare October 4, 2026 07:42
@Jutho

Jutho commented Oct 4, 2026

Copy link
Copy Markdown
Member

Ok, I did indeed not consider that ind is a Vector{Int} so often.

Comment thread src/pullbacks/svd.jl Outdated
Comment thread src/pullbacks/svd.jl Outdated
Comment thread src/pullbacks/svd.jl Outdated
@Jutho

Jutho commented Oct 5, 2026

Copy link
Copy Markdown
Member

Ok, I've pushed some changes. @leburgel, time for you or 🤖 to check whether this is still correct (and of course also CI).

@leburgel

leburgel commented Oct 5, 2026

Copy link
Copy Markdown
Member Author

I pushed ba679e9 on top of your typo fixes. Apart from the unpacking of ΔDV in eigh_pullback! (it now takes ind₀ from check_and_prepare_eigh_cotangents) and VᴴΔV₀' → VᴴΔV₁₀', two things gave wrong results:

  • check_and_prepare_svd_cotangents returned aUᴴΔU₁₀ instead of aUᴴΔAV₁₀, the antihermitian part of UᴴΔAV; the call sites now use that name.
  • The columns ind₀ of UᴴΔAV are hUᴴΔAV₁₀ .+ aUᴴΔAV₁₀, without adjoints; only the rows are hUᴴΔAV₁₀' .- aUᴴΔAV₁₀'.

I also rewrote the comment on the two blocks in svd_pullback!, which still described UᴴΔAVₖ and UᴴΔAVʳ; feel free to reword it.

Two changes in behaviour compared with main, which I left as they are: cotangents beyond the rank in the rank-deficient case are no longer checked to be zero, and for svd_full with ind = Colon() and rank below min(m, n) they now throw the ArgumentError. The tests don't cover the latter. Also, ΔU[:, J] and view(ΔS, J) index with a BitVector, which may need findall(J) on GPU.

@Jutho

Jutho commented Oct 5, 2026

Copy link
Copy Markdown
Member

Thanks. The aUᴴΔU -> aUᴴΔAV bug makes me wonder if this was something that VSCode autocorrect or copilot changed behind my back.

I'll look into the bitvector indexing next, because it is clearly causing issues on the GPU / buildkite CI.

@leburgel

leburgel commented Oct 6, 2026

Copy link
Copy Markdown
Member Author

I pushed 2e540de for the remaining Buildkite failure (svd_full in cuda / mooncake). For svd_full, ΔS = diagview(ΔSmat) is itself a view of the dense ΔSmat, so view(ΔS, J₁) on line 93 was a view of a view with a CPU Vector{Int} index, which cannot be passed to the GPU kernel; for svd_compact, ΔS is the vector of a Diagonal, which is why only svd_full failed. It is now ΔS[J₁], as S[ind₀] for S₀. I did the same for view(ΔS, J₂) on line 94 and for S₀ = view(S, ind₀) in svd_pullback!, which is used in the broadcasts with ΔU₊ and ΔV₊ᴴ, and applied Runic's 0:(length(ind₀) - 1) on line 93.

@leburgel
leburgel force-pushed the lb/full_pullback_kept_columns branch from 2e540de to 2e4ec15 Compare October 6, 2026 15:44
Comment thread src/pullbacks/svd.jl
leburgel and others added 6 commits October 7, 2026 06:57
@leburgel
leburgel force-pushed the lb/full_pullback_kept_columns branch from 2e4ec15 to 8051d01 Compare October 7, 2026 04:59
@leburgel
leburgel enabled auto-merge (squash) October 7, 2026 05:03
@leburgel

leburgel commented Oct 7, 2026

Copy link
Copy Markdown
Member Author

Buildkite shows a failure in the eig tests, but I don't think this should touches anything directly eig-related. Could the failure be unrelated, or does this still break something after all @Jutho?

@Jutho

Jutho commented Oct 7, 2026

Copy link
Copy Markdown
Member

No, the eig failures seem to be random accuracy glitches with Float32 precision. Not sure how this can be so undeterministic (maybe spooky particles hitting the GPU).

Screenshot 2026-10-07 at 16 37 02

@leburgel
leburgel merged commit 599babb into main Oct 8, 2026
48 of 50 checks passed
@leburgel
leburgel deleted the lb/full_pullback_kept_columns branch October 8, 2026 04:39
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