Skip to content

Sum only one Neumann series in svd_trunc_pullback! - #287

Open
leburgel wants to merge 4 commits into
mainfrom
lb/onesided_svd_trunc_pullback
Open

leburgel wants to merge 4 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 90.00000% with 1 line in your changes missing coverage. Please review.

Files with missing lines Patch % Lines
src/pullbacks/svd.jl 90.00% 1 Missing ⚠️
Files with missing lines Coverage Δ
src/pullbacks/svd.jl 93.22% <90.00%> (-0.14%) ⬇️

... and 1 file 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
S⁻¹ = minS ./ S
# sum the series on the smaller side only, the other side follows from Yᴴ = Y₀ᴴ + S⁻¹ X' AP;
# for m > n, work with the adjoint problem, which swaps X and Yᴴ
m > n && ((AP, X₀, Y₀ᴴ) = (AP', Y₀ᴴ', X₀'))

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 looks like it results in some type instability, but I'm not sure if the compiler just union-splits this correctly.
In any case, it might be better to make this a bit more explicit and manually write out the two cases, either duplicating the code or introducing a function barrier?

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.

I would also prefer to see this separated in the two cases 😄 .

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 separated the two cases, and I split off the actual doubling iteration in a separate _smith_iteration! routine that is now used in both the SVD and eigh pullbacks. I was told the doubling method is usually referred to as a "squared Smith iteration", hence the method name. The warning inside still says "Sylvester iteration", which I thought was fine since it makes it more clear we're solving a Sylvester equation.

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.

Comment thread src/common/pullbacks.jl
is_leading_index(ind::AbstractVector, p::Int) = length(ind) == p && all(ind .== 1:p)

"""
_smith_iteration!(X, Xₙ, G, w, atol, maxiter)

@Jutho Jutho Oct 2, 2026 •

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.

According to

https://www.sciencedirect.com/science/article/pii/S0893965909000263

this is the "Smith accelerative iteration", for which they also refer to this original Smith paper:

https://www.jstor.org/stable/pdf/2099416.pdf

Could you maybe add the reference and also change the name to

accelerative_smith_iteration!

I don't think the leading underscore is necessary; this can be a useful method to be used for other things as well.

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.

4 participants