From a50dbe4376601a647176e1d5f59742aa900bdc23 Mon Sep 17 00:00:00 2001 From: leburgel Date: Wed, 23 Sep 2026 12:09:10 +0200 Subject: [PATCH 1/4] Fix overflow in `eigh_trunc_pullback!` --- src/pullbacks/eigh.jl | 11 +-- test/common/eigh_trunc_pullback_overflow.jl | 74 +++++++++++++++++++++ 2 files changed, 81 insertions(+), 4 deletions(-) create mode 100644 test/common/eigh_trunc_pullback_overflow.jl diff --git a/src/pullbacks/eigh.jl b/src/pullbacks/eigh.jl index cc87ea484..a645e9d2d 100755 --- a/src/pullbacks/eigh.jl +++ b/src/pullbacks/eigh.jl @@ -153,9 +153,11 @@ function eigh_trunc_pullback!( if !iszerotangent(ΔV₊) X₀ = rdiv!(ΔV₊, Diagonal(D)) AP = mul!(copy(A), V * Dmat, V', -1, 1) - dabsmax = maximum(abs, D) - AP ./= dabsmax - D⁻¹ = dabsmax ./ D + # Normalize by the smallest retained |eigenvalue|, as `svd_trunc_pullback!` does + # with `S[end]`. That caps `max|D⁻¹|` at 1, so squaring can only shrink it. + dabsmin = minimum(abs, D) + AP ./= dabsmin + D⁻¹ = dabsmin ./ D X₁ = rmul!(AP * X₀, Diagonal(D⁻¹)) X₁ .+= X₀ Xₖ, Xₖ₊₁ = X₁, X₀ @@ -182,7 +184,8 @@ function eigh_trunc_pullback!( # take the Hermitian part, and cannot apply project_hermitian! to # the current contents of ΔA # TODO: add an `add_project_hermitian!` - ΔA′ = project_hermitian!(mul!(AP, Z, V', 1, 1)) # recycle AP + # recycle AP's storage, but overwrite it: the loop leaves APₖ in that buffer + ΔA′ = project_hermitian!(mul!(AP, Z, V', 1, 0)) ΔA .+= ΔA′ else # in this case, Z * V' is automatically Hermitian, so we can directly add it to ΔA diff --git a/test/common/eigh_trunc_pullback_overflow.jl b/test/common/eigh_trunc_pullback_overflow.jl new file mode 100644 index 000000000..e6a5e118a --- /dev/null +++ b/test/common/eigh_trunc_pullback_overflow.jl @@ -0,0 +1,74 @@ +using MatrixAlgebraKit +using MatrixAlgebraKit: eigh_pullback!, eigh_trunc_pullback!, diagview +using Test +using StableRNGs +using LinearAlgebra: Diagonal, norm, qr + +# The Sylvester doubling iteration in `eigh_trunc_pullback!` squares `APₖ` and `D⁻¹ₖ` +# separately although only their product is used. The product is bounded by +# ρ = |λ_discarded|max / |λ_retained|min < 1, but the factors are not, so if the series +# needs more doublings to converge than it takes them to leave the floating-point range, +# `APₖ` underflows to 0 while `D⁻¹ₖ` overflows to Inf and the next iterate is NaN. +# +# That needs ρ close to 1 — a well-separated truncation converges in a few doublings — so +# the spectrum below is chosen with +# +# |λ|max / |λ_kept|min = 30 -> D⁻¹^(2^k) overflows Float64 at 2^k > 208 (k = 8) +# ρ = 0.95 -> ρ^m < 1e-13 needs m > 583 (k = 10) + +@testset "eigh_trunc_pullback! with rho close to 1 ($T)" for T in (Float64, ComplexF64) + rng = StableRNG(12345) + n, p = 24, 6 + λ = vcat([30.0, -30.0, 5.0, -5.0, 1.0, -1.0], collect(range(0.95, 0.1; length = n - p))) + Q = Matrix(qr(randn(rng, T, n, n)).Q) + A = Q * Diagonal(T.(λ)) * Q' + A = (A + A') / 2 + + V, D = Q[:, 1:p], real(T).(λ[1:p]) # truncated factors taken straight from the + Dmat = Diagonal(D) # construction: no phase/ordering ambiguity + @test 0.9 < maximum(abs, λ[(p + 1):end]) / minimum(abs, D) < 1 + + ΔD = randn(rng, real(T), p) + ΔV = randn(rng, T, n, p) + ΔV .-= V * Diagonal(diagview(V' * ΔV)) # drop the gauge-sensitive component + + ΔA = eigh_trunc_pullback!(zeros(T, n, n), A, (Dmat, V), (Diagonal(ΔD), copy(ΔV))) + @test all(isfinite, ΔA) + @test ΔA ≈ ΔA' + + # reference: the same cotangents through the full pullback, which has no series in it + ΔA_full = eigh_pullback!( + zeros(T, n, n), A, (Diagonal(real(T).(λ)), Q), + (Diagonal(vcat(ΔD, zeros(real(T), n - p))), hcat(copy(ΔV), zeros(T, n, n - p))) + ) + @test ΔA ≈ ΔA_full rtol = 1.0e-8 +end + +# The final assembly recycles the `AP` buffer with `mul!(AP, Z, V', 1, α)`. `APₖ₊₁` aliases +# that buffer, so it must be OVERWRITTEN (α = 0), not accumulated into. With a well +# separated truncation the loop converges at k = 1, so no squaring has been written yet and +# the buffer still holds the original `AP`; accumulating adds a spurious term of relative +# size |λ_discarded|max / |λ|max. +@testset "eigh_trunc_pullback! does not accumulate into the recycled AP buffer ($T)" for T in + (Float64, ComplexF64) + + rng = StableRNG(12345) + n, p = 24, 6 + λ = vcat([30.0, -30.0, 5.0, -5.0, 1.0, -1.0], collect(range(1.0e-6, 1.0e-7; length = n - p))) + Q = Matrix(qr(randn(rng, T, n, n)).Q) + A = Q * Diagonal(T.(λ)) * Q' + A = (A + A') / 2 + + V, D = Q[:, 1:p], real(T).(λ[1:p]) + ΔD = randn(rng, real(T), p) + ΔV = randn(rng, T, n, p) + ΔV .-= V * Diagonal(diagview(V' * ΔV)) + + ΔA = eigh_trunc_pullback!(zeros(T, n, n), A, (Diagonal(D), V), (Diagonal(ΔD), copy(ΔV))) + ΔA_full = eigh_pullback!( + zeros(T, n, n), A, (Diagonal(real(T).(λ)), Q), + (Diagonal(vcat(ΔD, zeros(real(T), n - p))), hcat(copy(ΔV), zeros(T, n, n - p))) + ) + # the spurious term would be ~3e-8 relative, far above this + @test ΔA ≈ ΔA_full rtol = 1.0e-10 +end From 084a58a2c35ff4796a58f3c26f73cd00a1918f03 Mon Sep 17 00:00:00 2001 From: leburgel Date: Wed, 23 Sep 2026 14:59:18 +0200 Subject: [PATCH 2/4] Compress comment --- test/common/eigh_trunc_pullback_overflow.jl | 20 ++++++++------------ 1 file changed, 8 insertions(+), 12 deletions(-) diff --git a/test/common/eigh_trunc_pullback_overflow.jl b/test/common/eigh_trunc_pullback_overflow.jl index e6a5e118a..b31205c3d 100644 --- a/test/common/eigh_trunc_pullback_overflow.jl +++ b/test/common/eigh_trunc_pullback_overflow.jl @@ -4,21 +4,17 @@ using Test using StableRNGs using LinearAlgebra: Diagonal, norm, qr -# The Sylvester doubling iteration in `eigh_trunc_pullback!` squares `APₖ` and `D⁻¹ₖ` -# separately although only their product is used. The product is bounded by -# ρ = |λ_discarded|max / |λ_retained|min < 1, but the factors are not, so if the series -# needs more doublings to converge than it takes them to leave the floating-point range, -# `APₖ` underflows to 0 while `D⁻¹ₖ` overflows to Inf and the next iterate is NaN. -# -# That needs ρ close to 1 — a well-separated truncation converges in a few doublings — so -# the spectrum below is chosen with -# -# |λ|max / |λ_kept|min = 30 -> D⁻¹^(2^k) overflows Float64 at 2^k > 208 (k = 8) -# ρ = 0.95 -> ρ^m < 1e-13 needs m > 583 (k = 10) +# Regression test for the Sylvester doubling iteration in `eigh_trunc_pullback!`, +# which previously gave rise to overflow and underflow issues when the gap between the +# absolute values of the largest discarded and the smallest retained eigenvalues +# was small, but the gap between the absolute values of the largest and smallest retained +# eigenvalues was large. -@testset "eigh_trunc_pullback! with rho close to 1 ($T)" for T in (Float64, ComplexF64) +@testset "Regression test for eigh_trunc_pullback! with rho close to 1 ($T)" for T in (Float64, ComplexF64) rng = StableRNG(12345) n, p = 24, 6 + # construct artifical spectrum with |λ_discarded|max / |λ_kept|min close to 1, but + # |λ_kept|max / |λ_kept|min large λ = vcat([30.0, -30.0, 5.0, -5.0, 1.0, -1.0], collect(range(0.95, 0.1; length = n - p))) Q = Matrix(qr(randn(rng, T, n, n)).Q) A = Q * Diagonal(T.(λ)) * Q' From 667321ef1fb7e65ffbf9aa2b44fb5e9a957dd534 Mon Sep 17 00:00:00 2001 From: Lander Burgelman <39218680+leburgel@users.noreply.github.com> Date: Wed, 23 Sep 2026 15:35:55 +0200 Subject: [PATCH 3/4] Update src/pullbacks/eigh.jl Co-authored-by: Jutho --- src/pullbacks/eigh.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/pullbacks/eigh.jl b/src/pullbacks/eigh.jl index a645e9d2d..713c04b5c 100755 --- a/src/pullbacks/eigh.jl +++ b/src/pullbacks/eigh.jl @@ -185,7 +185,7 @@ function eigh_trunc_pullback!( # the current contents of ΔA # TODO: add an `add_project_hermitian!` # recycle AP's storage, but overwrite it: the loop leaves APₖ in that buffer - ΔA′ = project_hermitian!(mul!(AP, Z, V', 1, 0)) + ΔA′ = project_hermitian!(mul!(AP, Z, V')) ΔA .+= ΔA′ else # in this case, Z * V' is automatically Hermitian, so we can directly add it to ΔA From 5035dad90ed6efc7d05b5b7506059418e0c66ea7 Mon Sep 17 00:00:00 2001 From: leburgel Date: Wed, 23 Sep 2026 15:39:26 +0200 Subject: [PATCH 4/4] Remove contrived regression test, update the name of the other --- test/common/eigh_trunc_pullback_overflow.jl | 31 +-------------------- 1 file changed, 1 insertion(+), 30 deletions(-) diff --git a/test/common/eigh_trunc_pullback_overflow.jl b/test/common/eigh_trunc_pullback_overflow.jl index b31205c3d..ad964c274 100644 --- a/test/common/eigh_trunc_pullback_overflow.jl +++ b/test/common/eigh_trunc_pullback_overflow.jl @@ -10,7 +10,7 @@ using LinearAlgebra: Diagonal, norm, qr # was small, but the gap between the absolute values of the largest and smallest retained # eigenvalues was large. -@testset "Regression test for eigh_trunc_pullback! with rho close to 1 ($T)" for T in (Float64, ComplexF64) +@testset "Regression test for overflow in eigh_trunc_pullback! ($T)" for T in (Float64, ComplexF64) rng = StableRNG(12345) n, p = 24, 6 # construct artifical spectrum with |λ_discarded|max / |λ_kept|min close to 1, but @@ -39,32 +39,3 @@ using LinearAlgebra: Diagonal, norm, qr ) @test ΔA ≈ ΔA_full rtol = 1.0e-8 end - -# The final assembly recycles the `AP` buffer with `mul!(AP, Z, V', 1, α)`. `APₖ₊₁` aliases -# that buffer, so it must be OVERWRITTEN (α = 0), not accumulated into. With a well -# separated truncation the loop converges at k = 1, so no squaring has been written yet and -# the buffer still holds the original `AP`; accumulating adds a spurious term of relative -# size |λ_discarded|max / |λ|max. -@testset "eigh_trunc_pullback! does not accumulate into the recycled AP buffer ($T)" for T in - (Float64, ComplexF64) - - rng = StableRNG(12345) - n, p = 24, 6 - λ = vcat([30.0, -30.0, 5.0, -5.0, 1.0, -1.0], collect(range(1.0e-6, 1.0e-7; length = n - p))) - Q = Matrix(qr(randn(rng, T, n, n)).Q) - A = Q * Diagonal(T.(λ)) * Q' - A = (A + A') / 2 - - V, D = Q[:, 1:p], real(T).(λ[1:p]) - ΔD = randn(rng, real(T), p) - ΔV = randn(rng, T, n, p) - ΔV .-= V * Diagonal(diagview(V' * ΔV)) - - ΔA = eigh_trunc_pullback!(zeros(T, n, n), A, (Diagonal(D), V), (Diagonal(ΔD), copy(ΔV))) - ΔA_full = eigh_pullback!( - zeros(T, n, n), A, (Diagonal(real(T).(λ)), Q), - (Diagonal(vcat(ΔD, zeros(real(T), n - p))), hcat(copy(ΔV), zeros(T, n, n - p))) - ) - # the spurious term would be ~3e-8 relative, far above this - @test ΔA ≈ ΔA_full rtol = 1.0e-10 -end