Skip to content

Solve the svdsolve and eigsolve GMRES/BiCGStab pullbacks on the complement of the kept vectors - #173

Open
leburgel wants to merge 3 commits into
Jutho:masterfrom
leburgel:lb/gmres_pullback_complement
Open

leburgel wants to merge 3 commits into
Jutho:masterfrom
leburgel:lb/gmres_pullback_complement

Conversation

@leburgel

Copy link
Copy Markdown
Contributor

With alg_rrule = GMRES() or BiCGStab(), the svdsolve and eigsolve reverse rules solve one linear problem per kept vector on the whole space, projecting out only that vector (bordering on it, for eigsolve). The other kept vectors stay in the problem, where the operator has eigenvalues ±σⱼ - σᵢ (svdsolve) or λ̄ⱼ - λ̄ᵢ (eigsolve): small when kept values are close, and of both signs. For more than a few kept vectors, restarted GMRES then does not converge and the gradient is wrong (reproducers below).

This PR takes the kept vectors out of the linear problem:

  • svdsolve: the linear problem is posed on the orthogonal complement of all kept singular vectors, where its spectrum σᵢ ± sⱼ is positive for discarded sⱼ < σᵢ, and the components along the other kept vectors follow in closed form from their 2 × 2 blocks (with safe_inv for degenerate σⱼ, as in the Arnoldi rule).
  • eigsolve: the orthogonal complement of the kept eigenvectors V is invariant under fᴴ, also for non-Hermitian f, with only the discarded eigenvalues there. Writing w = V a + y with y ⊥ V, the kept block (V's Gram matrix, as in the Arnoldi rule) follows in closed form and only y needs a linear solve, which no longer needs the border. This covers Lanczos and Arnoldi alike.

The gauge and convergence warnings are unchanged; eigsolve's "returns unexpected result" check is removed, since the border multiplier is now zero by construction. In five svdsolve tests at STARTSTOP_LEVEL the cotangents lie in the span of the kept vectors, so the linear problem on the complement is trivial and GMRES logs one message per vector instead of two; their expected log patterns are updated. test/ad/svdsolve.jl (220) and test/ad/eigsolve.jl (618) pass.

The reproducers below (N = 400, 10 kept vectors, GMRES at tol = 1e-10) compare the gradient of a gauge-invariant loss with the one through the dense svd / eigen:

master this PR
svdsolve error 4.4e-2, 5.8 s 2.9e-11, 0.46 s
eigsolve, Hermitian (Lanczos) error 3.7e-6, 1.7 s 1.6e-11, 0.21 s
eigsolve, non-Hermitian (Arnoldi) error 0.83, 3.3 s 5.6e-12, 0.76 s

On master each case also warns that some of the cotangent linear problems did not converge (shown once each below; the scripts evaluate the gradient twice).

svdsolve
using KrylovKit, LinearAlgebra, Random, Zygote

# the "Large svdsolve AD test" matrix (test/ad/svdsolve.jl) at N = 400, with k = 10 triplets
Random.seed!(1)
N, n, k = 400, 133, 10
A = rand(N, N + n) .- 1 / 2
A = I[1:N, 1:(N + n)] - (9 // 10) * A / maximum(svdvals(A))
x₀ = normalize!(randn(N))
c = randn(N)

# a loss that depends on the singular values and (gauge-invariantly) on the singular vectors
function loss(A, alg_rrule)
    vals, lvecs, = svdsolve(A, x₀, k, :LR, GKL(; tol = 1.0e-12, krylovdim = 40); alg_rrule)
    return sum(vals[1:k]) + sum(abs2(dot(c, lvecs[i])) for i in 1:k)
end
loss_exact(A) = (F = svd(A); sum(F.S[1:k]) + sum(abs2(dot(c, F.U[:, i])) for i in 1:k))

g_exact = only(Zygote.gradient(loss_exact, A))
alg_rrule = GMRES(; tol = 1.0e-10, krylovdim = 30, maxiter = 200, verbosity = 0)
g = only(Zygote.gradient(A -> loss(A, alg_rrule), A))
t = @elapsed Zygote.gradient(A -> loss(A, alg_rrule), A)
@show norm(g - g_exact) / norm(g_exact) t

On master:

Warning: `svdsolve` cotangent linear problem (5) did not converge, whereas the primal eigenvalue problem did: normres = 1.3480795515355712e-6
Warning: `svdsolve` cotangent linear problem (6) did not converge, whereas the primal eigenvalue problem did: normres = 0.0021100224895391363
Warning: `svdsolve` cotangent linear problem (7) did not converge, whereas the primal eigenvalue problem did: normres = 8.504932491666868e-10
Warning: `svdsolve` cotangent linear problem (8) did not converge, whereas the primal eigenvalue problem did: normres = 0.0329010861512903
Warning: `svdsolve` cotangent linear problem (9) did not converge, whereas the primal eigenvalue problem did: normres = 0.05149676946753317
Warning: `svdsolve` cotangent linear problem (10) did not converge, whereas the primal eigenvalue problem did: normres = 4.5350249179202516e-7
norm(g - g_exact) / norm(g_exact) = 0.043915707034682026
t = 5.753016758

With this PR:

norm(g - g_exact) / norm(g_exact) = 2.890852787628183e-11
t = 0.461598365
eigsolve, Hermitian
using KrylovKit, LinearAlgebra, Random, Zygote

# a Hermitian matrix like the "Large Hermitian eigsolve AD test" one (test/ad/eigsolve.jl) at
# N = 400, with k = 10 eigenpairs
Random.seed!(1)
N, k = 400, 10
A = rand(N, N) .- 1 / 2
A = (A + A') / 2
A = I - (9 // 10) * A / maximum(abs, eigvals(Hermitian(A)))
x₀ = normalize!(randn(N))
c = randn(N)

# a loss that depends on the eigenvalues and (gauge-invariantly) on the eigenvectors
function loss(A, alg_rrule)
    vals, vecs, = eigsolve(A, x₀, k, :LR, Lanczos(; tol = 1.0e-12, krylovdim = 40); alg_rrule)
    return sum(vals[1:k]) + sum(abs2(dot(c, vecs[i])) for i in 1:k)
end
function loss_exact(A)
    F = eigen(Hermitian(A))
    idx = sortperm(F.values; rev = true)[1:k]
    return sum(F.values[idx]) + sum(abs2(dot(c, F.vectors[:, i])) for i in idx)
end

g_exact = only(Zygote.gradient(loss_exact, A))
alg_rrule = GMRES(; tol = 1.0e-10, krylovdim = 30, maxiter = 200, verbosity = 0)
g = only(Zygote.gradient(A -> loss(A, alg_rrule), A))
t = @elapsed Zygote.gradient(A -> loss(A, alg_rrule), A)
herm(X) = (X + X') / 2   # only the Hermitian part of the gradient is determined
@show norm(herm(g) - herm(g_exact)) / norm(herm(g_exact)) t

On master:

Warning: `eigsolve` cotangent linear problem (7) did not converge, whereas the primal eigenvalue problem did: normres = 3.2945955451053963e-6
Warning: `eigsolve` cotangent linear problem (9) did not converge, whereas the primal eigenvalue problem did: normres = 3.181370121561197e-6
Warning: `eigsolve` cotangent linear problem (10) did not converge, whereas the primal eigenvalue problem did: normres = 1.5697001012058855e-7
norm(herm(g) - herm(g_exact)) / norm(herm(g_exact)) = 3.730485350071469e-6
t = 1.702511799

With this PR:

norm(herm(g) - herm(g_exact)) / norm(herm(g_exact)) = 1.5878889551606496e-11
t = 0.213437454
eigsolve, non-Hermitian
using KrylovKit, LinearAlgebra, Random, Zygote

# a non-Hermitian matrix like the "Large eigsolve AD test" one (test/ad/eigsolve.jl) at N = 400,
# with k ≥ 10 eigenpairs, k chosen so that no complex conjugate pair is split
Random.seed!(1)
N = 400
A = rand(N, N) .- 1 / 2
A = I - (9 // 10) * A / maximum(abs, eigvals(A))
λs = sort(eigvals(A); by = real, rev = true)
k = findfirst(j -> j ≥ 10 && !(real(λs[j]) ≈ real(λs[j + 1])), eachindex(λs))
x₀ = normalize!(randn(N))
c = randn(N)

# a loss that depends on the eigenvalues and (gauge-invariantly) on the eigenvectors
function loss(A, alg_rrule)
    vals, vecs, = eigsolve(A, x₀, k, :LR, Arnoldi(; tol = 1.0e-12, krylovdim = 60); alg_rrule)
    return sum(real(vals[1:k])) + sum(abs2(dot(c, vecs[i])) for i in 1:k)
end
function loss_exact(A)
    F = eigen(A)
    idx = sortperm(F.values; by = real, rev = true)[1:k]
    return sum(real(F.values[idx])) + sum(abs2(dot(c, F.vectors[:, i])) for i in idx)
end

g_exact = only(Zygote.gradient(loss_exact, A))
alg_rrule = GMRES(; tol = 1.0e-10, krylovdim = 30, maxiter = 200, verbosity = 0)
g = only(Zygote.gradient(A -> loss(A, alg_rrule), A))
t = @elapsed Zygote.gradient(A -> loss(A, alg_rrule), A)
@show k norm(g - g_exact) / norm(g_exact) t

On master:

Warning: `eigsolve` cotangent linear problem (6) did not converge, whereas the primal eigenvalue problem did: normres = 11.657081309720454
Warning: `eigsolve` cotangent linear problem (7) did not converge, whereas the primal eigenvalue problem did: normres = 11.657081309720454
Warning: `eigsolve` cotangent linear problem (10) did not converge, whereas the primal eigenvalue problem did: normres = 11.471387620287413
k = 10
norm(g - g_exact) / norm(g_exact) = 0.8315597072059524
t = 3.279510926

With this PR:

k = 10
norm(g - g_exact) / norm(g_exact) = 5.553040586239211e-12
t = 0.762937438

@codecov

codecov Bot commented Sep 30, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 86.70%. Comparing base (4808454) to head (d6eadf8).
⚠️ Report is 1 commits behind head on master.

Additional details and impacted files
@@            Coverage Diff             @@
##           master     #173      +/-   ##
==========================================
- Coverage   88.47%   86.70%   -1.78%     
==========================================
  Files          36       36              
  Lines        3965     3933      -32     
==========================================
- Hits         3508     3410      -98     
- Misses        457      523      +66     

☔ View full report in Codecov by Harness.
📢 Have feedback on the report? Share it here.

🚀 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