Skip to content

Sum only one Neumann series in svd_trunc_pullback! - #287

Merged
lkdvos merged 7 commits into
mainfrom
lb/onesided_svd_trunc_pullback
Oct 2, 2026
Merged

lkdvos merged 7 commits into
mainfrom
lb/onesided_svd_trunc_pullback

Conversation

@leburgel

Copy link
Copy Markdown
Member

svd_trunc_pullback! sums the Neumann series of both complement equations by doubling, for X with AP * AP' (m × m) and for Yᴴ with AP' * AP (n × n). The two are not independent: Yᴴ = Y₀ᴴ + S⁻¹ X' AP. This PR sums the series on the smaller side only and gets the other one from that single product (for m > n it works with the adjoint problem, which swaps X and Yᴴ). Each doubling step then squares one Gram matrix instead of two, and the stopping criterion checks the summed side only.

It's the same series with the same stopping criterion: on the cases below the result agrees with main to at most 2.4e-15 (relative), and it is 2.0–2.6× faster for the square and small cases and 3.5–5.4× for the large rectangular ones, where the larger of the two Gram matrices drops out (except 1.6× for real 1400×1000; laptop timings, minimum of 5).

Benchmark (time and error against the full-spectrum `svd_pullback!`)
using LinearAlgebra, Random, Printf
using MatrixAlgebraKit
using MatrixAlgebraKit: svd_pullback!, svd_trunc_pullback!, remove_svd_gauge_dependence!, diagview

# time and accuracy of svd_trunc_pullback! against the full-spectrum svd_pullback!, with
# cotangents built as in the package tests (test/testsuite/ad_utils.jl)
function case(T, m, n, p)
    rng = Xoshiro(1)
    A = randn(rng, T, m, n)
    U, S, Vᴴ = svd_compact(A)
    ΔU, ΔVᴴ = remove_svd_gauge_dependence!(randn(rng, T, size(U)), randn(rng, T, size(Vᴴ)), U, S, Vᴴ)
    ΔS = Diagonal(randn(rng, real(T), size(S, 1)))
    ind = 1:p
    trunc = (U[:, ind], Diagonal(diagview(S)[ind]), Vᴴ[ind, :])
    Δtrunc = (ΔU[:, ind], Diagonal(diagview(ΔS)[ind]), ΔVᴴ[ind, :])
    ref = svd_pullback!(zero(A), A, (U, S, Vᴴ), Δtrunc, ind)
    g = svd_trunc_pullback!(zero(A), A, trunc, Δtrunc)
    t = minimum(@elapsed(svd_trunc_pullback!(zero(A), A, trunc, Δtrunc)) for _ in 1:5)
    @printf("%-10s %4d×%-4d p=%-3d  time %.3e s  error %.1e\n", T, m, n, p, t, norm(g - ref) / norm(ref))
end

BLAS.set_num_threads(4)
for T in (Float64, ComplexF64), (m, n, p) in ((19, 17, 5), (19, 23, 5), (400, 400, 40), (1000, 1000, 100), (1000, 1400, 100), (1400, 1000, 100))
    case(T, m, n, p)
end

On main:

Float64      19×17   p=5    time 3.357e-05 s  error 1.0e-15
Float64      19×23   p=5    time 4.616e-05 s  error 3.4e-15
Float64     400×400  p=40   time 1.224e-01 s  error 3.3e-15
Float64    1000×1000 p=100  time 1.565e+00 s  error 2.2e-15
Float64    1000×1400 p=100  time 2.439e+00 s  error 7.2e-13
Float64    1400×1000 p=100  time 1.155e+00 s  error 2.9e-13
ComplexF64   19×17   p=5    time 6.484e-05 s  error 7.9e-16
ComplexF64   19×23   p=5    time 8.254e-05 s  error 2.6e-15
ComplexF64  400×400  p=40   time 3.139e-01 s  error 4.2e-15
ComplexF64 1000×1000 p=100  time 2.235e+00 s  error 2.5e-14
ComplexF64 1000×1400 p=100  time 6.329e+00 s  error 1.8e-12
ComplexF64 1400×1000 p=100  time 5.828e+00 s  error 6.8e-14

With this PR:

Float64      19×17   p=5    time 1.488e-05 s  error 9.9e-16
Float64      19×23   p=5    time 2.085e-05 s  error 1.4e-15
Float64     400×400  p=40   time 4.795e-02 s  error 3.3e-15
Float64    1000×1000 p=100  time 7.939e-01 s  error 2.3e-15
Float64    1000×1400 p=100  time 6.937e-01 s  error 7.2e-13
Float64    1400×1000 p=100  time 7.273e-01 s  error 2.9e-13
ComplexF64   19×17   p=5    time 2.847e-05 s  error 7.8e-16
ComplexF64   19×23   p=5    time 3.695e-05 s  error 2.7e-15
ComplexF64  400×400  p=40   time 1.444e-01 s  error 4.2e-15
ComplexF64 1000×1000 p=100  time 1.112e+00 s  error 2.5e-14
ComplexF64 1000×1400 p=100  time 1.332e+00 s  error 1.8e-12
ComplexF64 1400×1000 p=100  time 1.083e+00 s  error 6.8e-14

@leburgel
leburgel requested a review from Jutho September 30, 2026 14:23
@codecov

codecov Bot commented Oct 1, 2026 •

Copy link
Copy Markdown

Codecov Report

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

Files with missing lines Patch % Lines
src/common/pullbacks.jl 88.88% 2 Missing ⚠️
Files with missing lines Coverage Δ
src/pullbacks/eigh.jl 86.07% <100.00%> (-0.07%) ⬇️
src/pullbacks/svd.jl 94.47% <100.00%> (+1.00%) ⬆️
src/common/pullbacks.jl 92.85% <88.88%> (-7.15%) ⬇️

... and 14 files with indirect coverage changes

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@leburgel
leburgel force-pushed the lb/onesided_svd_trunc_pullback branch from 63b2cb2 to fc70f25 Compare October 1, 2026 06:38
Comment thread src/pullbacks/svd.jl Outdated
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 1, 2026

Copy link
Copy Markdown
Member

Ok, very nice, and also somewhat trivial in hindsight. Very stupid of me to not spot this when I was deriving this.

@leburgel
leburgel force-pushed the lb/onesided_svd_trunc_pullback branch from 4dd40de to 90af904 Compare October 2, 2026 06:42
@leburgel
leburgel added this pull request to stack #295 October 2, 2026 07:01
Comment thread src/common/pullbacks.jl Outdated
Comment thread src/common/pullbacks.jl Outdated
Comment thread src/common/pullbacks.jl
leburgel and others added 2 commits October 2, 2026 15:51
Co-authored-by: Jutho <Jutho@users.noreply.github.com>
@lkdvos

lkdvos commented Oct 2, 2026

Copy link
Copy Markdown
Member

@leburgel I don't really know what this stacking feature is, but that seems to say that we can only merge this if both of these changes are approved. I feel like this one is probably self-contained enough to be merged as-is?

@leburgel

leburgel commented Oct 2, 2026 •

Copy link
Copy Markdown
Member Author

@leburgel I don't really know what this stacking feature is, but that seems to say that we can only merge this if both of these changes are approved. I feel like this one is probably self-contained enough to be merged as-is?

I think this one just says it can't be merged because the tests didn't complete yet. I can still bypass and force merge here, so I think it just obeys the same rules. The stack is only so the upper one only shows the relevant changes without targeting it on a branch that will be removed once the lower one is merged. I think for this PR this shouldn't change any normal behavior. First time trying this out though, so I could be wrong.

Do you want to force merge this or wait for the tests to complete?

@lkdvos

lkdvos commented Oct 2, 2026

Copy link
Copy Markdown
Member

Hmmm, guess I just got confused because it also does not give me the auto-merge option, and it seems to indicate that it wants to merge only if all PRs in the stack are ready. I'll keep an eye on this today and merge when tests complete, I don't think this is that urgent?

@leburgel

leburgel commented Oct 2, 2026

Copy link
Copy Markdown
Member Author

Not urgent no. I didn't notice the auto-merge option wasn't there, but it seems this is just a missing feature for stacks: github/gh-stack#239. If there's a problem with merging in the end, I'll un-stack.

@lkdvos
lkdvos merged commit a16a53f into main Oct 2, 2026
44 of 46 checks passed
@lkdvos
lkdvos deleted the lb/onesided_svd_trunc_pullback branch October 2, 2026 18:02
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.

4 participants