From 76dd6120c06540993f308e826b74756dce3023b4 Mon Sep 17 00:00:00 2001 From: leburgel Date: Wed, 30 Sep 2026 16:19:32 +0200 Subject: [PATCH 1/7] Sum only one Neumann series in `svd_trunc_pullback!` --- src/pullbacks/svd.jl | 26 ++++++++++++-------------- 1 file changed, 12 insertions(+), 14 deletions(-) diff --git a/src/pullbacks/svd.jl b/src/pullbacks/svd.jl index 1c894528f..74e00d8da 100755 --- a/src/pullbacks/svd.jl +++ b/src/pullbacks/svd.jl @@ -250,36 +250,34 @@ function svd_trunc_pullback!( minS = @view S[end:end] AP ./= minS 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₀')) X₁ = rmul!(AP * Y₀ᴴ', Diagonal(S⁻¹)) X₁ .+= X₀ - Y₁ᴴ = lmul!(Diagonal(S⁻¹), X₀' * AP) - Y₁ᴴ .+= Y₀ᴴ - Xₖ, Xₖ₊₁ = X₁, X₀ - Yₖᴴ, Yₖ₊₁ᴴ = Y₁ᴴ, Y₀ᴴ - APAᴴₖ, AᴴPAₖ = AP * AP', AP' * AP - APAᴴₖ₊₁, AᴴPAₖ₊₁ = zero(APAᴴₖ), zero(AᴴPAₖ) - S⁻¹ₖ, S⁻¹ₖ₊₁ = S⁻¹ .^ 2, S⁻¹ + Xₖ, Xₖ₊₁ = X₁, zero(X₁) + APAᴴₖ = AP * AP' + APAᴴₖ₊₁ = zero(APAᴴₖ) + S⁻¹ₖ, S⁻¹ₖ₊₁ = S⁻¹ .^ 2, zero(S⁻¹) for k in 1:maxiter Xₖ₊₁ = rmul!(mul!(Xₖ₊₁, APAᴴₖ, Xₖ), Diagonal(S⁻¹ₖ)) - Yₖ₊₁ᴴ = lmul!(Diagonal(S⁻¹ₖ), mul!(Yₖ₊₁ᴴ, Yₖᴴ, AᴴPAₖ)) - if norm(Xₖ₊₁, Inf) < degeneracy_atol && norm(Yₖ₊₁ᴴ, Inf) < degeneracy_atol + if norm(Xₖ₊₁, Inf) < degeneracy_atol break end Xₖ₊₁ .+= Xₖ - Yₖ₊₁ᴴ .+= Yₖᴴ if k == maxiter - @warn "Sylvester iteration did not converge after $k iterations, final norms of X: $(norm(Xₖ₊₁, Inf)), Yᴴ: $(norm(Yₖ₊₁ᴴ, Inf)))" + @warn "Sylvester iteration did not converge after $k iterations, final norm of X: $(norm(Xₖ₊₁, Inf)))" break end S⁻¹ₖ₊₁ .= S⁻¹ₖ .^ 2 APAᴴₖ₊₁ = mul!(APAᴴₖ₊₁, APAᴴₖ, APAᴴₖ) - AᴴPAₖ₊₁ = mul!(AᴴPAₖ₊₁, AᴴPAₖ, AᴴPAₖ) Xₖ, Xₖ₊₁ = Xₖ₊₁, Xₖ - Yₖᴴ, Yₖ₊₁ᴴ = Yₖ₊₁ᴴ, Yₖᴴ APAᴴₖ, APAᴴₖ₊₁ = APAᴴₖ₊₁, APAᴴₖ - AᴴPAₖ, AᴴPAₖ₊₁ = AᴴPAₖ₊₁, AᴴPAₖ S⁻¹ₖ, S⁻¹ₖ₊₁ = S⁻¹ₖ₊₁, S⁻¹ₖ end + Yₖᴴ = lmul!(Diagonal(S⁻¹), Xₖ' * AP) + Yₖᴴ .+= Y₀ᴴ + m > n && ((Xₖ, Yₖᴴ) = (Yₖᴴ', Xₖ')) ΔA = mul!(ΔA, Xₖ, Vᴴ, 1, 1) ΔA = mul!(ΔA, U, Yₖᴴ, 1, 1) end From 6ec9e0ad60b2f8496391ee6673c686b6de06e39c Mon Sep 17 00:00:00 2001 From: leburgel Date: Thu, 1 Oct 2026 08:56:09 +0200 Subject: [PATCH 2/7] Fix GPU safety --- src/pullbacks/svd.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/pullbacks/svd.jl b/src/pullbacks/svd.jl index 74e00d8da..d9778920e 100755 --- a/src/pullbacks/svd.jl +++ b/src/pullbacks/svd.jl @@ -266,7 +266,7 @@ function svd_trunc_pullback!( end Xₖ₊₁ .+= Xₖ if k == maxiter - @warn "Sylvester iteration did not converge after $k iterations, final norm of X: $(norm(Xₖ₊₁, Inf)))" + @warn "Sylvester iteration did not converge after $k iterations, final norm of X: $(maximum(abs, Xₖ₊₁))" break end S⁻¹ₖ₊₁ .= S⁻¹ₖ .^ 2 From 312c3430d91dbc4f852063d955f5b9f5717baf22 Mon Sep 17 00:00:00 2001 From: Lander Burgelman <39218680+leburgel@users.noreply.github.com> Date: Thu, 1 Oct 2026 19:50:37 +0200 Subject: [PATCH 3/7] Update src/pullbacks/svd.jl Co-authored-by: Jutho --- src/pullbacks/svd.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/pullbacks/svd.jl b/src/pullbacks/svd.jl index d9778920e..b8dc897d1 100755 --- a/src/pullbacks/svd.jl +++ b/src/pullbacks/svd.jl @@ -255,7 +255,7 @@ function svd_trunc_pullback!( m > n && ((AP, X₀, Y₀ᴴ) = (AP', Y₀ᴴ', X₀')) X₁ = rmul!(AP * Y₀ᴴ', Diagonal(S⁻¹)) X₁ .+= X₀ - Xₖ, Xₖ₊₁ = X₁, zero(X₁) + Xₖ, Xₖ₊₁ = X₁, X₀ APAᴴₖ = AP * AP' APAᴴₖ₊₁ = zero(APAᴴₖ) S⁻¹ₖ, S⁻¹ₖ₊₁ = S⁻¹ .^ 2, zero(S⁻¹) From 90af904d52cca8c446948ba6079c95056000430a Mon Sep 17 00:00:00 2001 From: leburgel Date: Fri, 2 Oct 2026 08:41:00 +0200 Subject: [PATCH 4/7] Split off reusable doubling iteration, split cases for rectangular SVD --- src/common/pullbacks.jl | 31 +++++++++++++++++++++++++++ src/pullbacks/eigh.jl | 26 +++-------------------- src/pullbacks/svd.jl | 47 ++++++++++++++++------------------------- 3 files changed, 52 insertions(+), 52 deletions(-) diff --git a/src/common/pullbacks.jl b/src/common/pullbacks.jl index 053db5335..60b254eca 100644 --- a/src/common/pullbacks.jl +++ b/src/common/pullbacks.jl @@ -35,3 +35,34 @@ iterating over `ind`, so that this also works for an `ind` that lives on a devic """ is_leading_index(ind::AbstractRange, p::Int) = ind == 1:p is_leading_index(ind::AbstractVector, p::Int) = length(ind) == p && all(ind .== 1:p) + +""" + _smith_iteration!(X, Xₙ, G, w, atol, maxiter) + +Solve `X = B + G * X * Diagonal(w)` by summing the Neumann series +`X = Σₖ Gᵏ * B * Diagonal(w)ᵏ` by doubling (Smith's method), i.e. by repeatedly adding +`G^(2ʲ) * X * Diagonal(w)^(2ʲ)` to `X` until the norm of that increment drops below `atol`, +for at most `maxiter` steps. + +On entry, `X` contains `B`, and it is overwritten with the result. `Xₙ` is used as a buffer, +and `G` and `w` are overwritten. It is assumed that `w` is normalized such that +`maximum(abs, w) == 1`, so that squaring it can only shrink it. +""" +function _smith_iteration!(X, Xₙ, G, w, atol, maxiter) + Gₙ = similar(G) + for k in 1:maxiter + Xₙ = rmul!(mul!(Xₙ, G, X), Diagonal(w)) + if maximum(abs, Xₙ) < atol + break + end + X .+= Xₙ + if k == maxiter + @warn "Sylvester iteration did not converge after $k iterations, final norm of X: $(maximum(abs, X))" + break + end + w .= w .^ 2 + Gₙ = mul!(Gₙ, G, G) + G, Gₙ = Gₙ, G + end + return X +end diff --git a/src/pullbacks/eigh.jl b/src/pullbacks/eigh.jl index 713c04b5c..eaafe820f 100755 --- a/src/pullbacks/eigh.jl +++ b/src/pullbacks/eigh.jl @@ -158,33 +158,13 @@ function eigh_trunc_pullback!( dabsmin = minimum(abs, D) AP ./= dabsmin D⁻¹ = dabsmin ./ D - X₁ = rmul!(AP * X₀, Diagonal(D⁻¹)) - X₁ .+= X₀ - Xₖ, Xₖ₊₁ = X₁, X₀ - APₖ, APₖ₊₁ = AP * AP, AP - D⁻¹ₖ, D⁻¹ₖ₊₁ = D⁻¹ .^ 2, D⁻¹ - for k in 1:maxiter - Xₖ₊₁ = rmul!(mul!(Xₖ₊₁, APₖ, Xₖ), Diagonal(D⁻¹ₖ)) - if norm(Xₖ₊₁, Inf) < degeneracy_atol - break - end - Xₖ₊₁ .+= Xₖ - if k == maxiter - @warn "Sylvester iteration did not converge after $k iterations, final norm of X: $(norm(Xₖ₊₁, Inf)))" - break - end - D⁻¹ₖ₊₁ .= D⁻¹ₖ .^ 2 - APₖ₊₁ = mul!(APₖ₊₁, APₖ, APₖ) - Xₖ, Xₖ₊₁ = Xₖ₊₁, Xₖ - APₖ, APₖ₊₁ = APₖ₊₁, APₖ - D⁻¹ₖ, D⁻¹ₖ₊₁ = D⁻¹ₖ₊₁, D⁻¹ₖ - end - Z .+= Xₖ + X = _smith_iteration!(X₀, similar(X₀), AP, D⁻¹, degeneracy_atol, maxiter) + Z .+= X # we cannot directly multiply Z * V' into ΔA, because we have to # take the Hermitian part, and cannot apply project_hermitian! to # 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 + # recycle AP's storage, but overwrite it: `_smith_iteration!` may leave a power of AP in it ΔA′ = project_hermitian!(mul!(AP, Z, V')) ΔA .+= ΔA′ else diff --git a/src/pullbacks/svd.jl b/src/pullbacks/svd.jl index b8dc897d1..d0883c65e 100755 --- a/src/pullbacks/svd.jl +++ b/src/pullbacks/svd.jl @@ -250,36 +250,25 @@ function svd_trunc_pullback!( minS = @view S[end:end] AP ./= minS 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₀')) - X₁ = rmul!(AP * Y₀ᴴ', Diagonal(S⁻¹)) - X₁ .+= X₀ - Xₖ, Xₖ₊₁ = X₁, X₀ - APAᴴₖ = AP * AP' - APAᴴₖ₊₁ = zero(APAᴴₖ) - S⁻¹ₖ, S⁻¹ₖ₊₁ = S⁻¹ .^ 2, zero(S⁻¹) - for k in 1:maxiter - Xₖ₊₁ = rmul!(mul!(Xₖ₊₁, APAᴴₖ, Xₖ), Diagonal(S⁻¹ₖ)) - if norm(Xₖ₊₁, Inf) < degeneracy_atol - break - end - Xₖ₊₁ .+= Xₖ - if k == maxiter - @warn "Sylvester iteration did not converge after $k iterations, final norm of X: $(maximum(abs, Xₖ₊₁))" - break - end - S⁻¹ₖ₊₁ .= S⁻¹ₖ .^ 2 - APAᴴₖ₊₁ = mul!(APAᴴₖ₊₁, APAᴴₖ, APAᴴₖ) - Xₖ, Xₖ₊₁ = Xₖ₊₁, Xₖ - APAᴴₖ, APAᴴₖ₊₁ = APAᴴₖ₊₁, APAᴴₖ - S⁻¹ₖ, S⁻¹ₖ₊₁ = S⁻¹ₖ₊₁, S⁻¹ₖ + # sum the series on the smaller side only, the other side follows from + # Yᴴ = Y₀ᴴ + S⁻¹ X' AP (m ≤ n) or X = X₀ + AP Y S⁻¹ (m > n) + if m ≤ n + X = rmul!(AP * Y₀ᴴ', Diagonal(S⁻¹)) + X .+= X₀ + X = _smith_iteration!(X, X₀, AP * AP', S⁻¹ .^ 2, degeneracy_atol, maxiter) # recycle X₀ + Yᴴ = lmul!(Diagonal(S⁻¹), X' * AP) + Yᴴ .+= Y₀ᴴ + ΔA = mul!(ΔA, X, Vᴴ, 1, 1) + ΔA = mul!(ΔA, U, Yᴴ, 1, 1) + else + Y = rmul!(AP' * X₀, Diagonal(S⁻¹)) + Y .+= Y₀ᴴ' + Y = _smith_iteration!(Y, similar(Y), AP' * AP, S⁻¹ .^ 2, degeneracy_atol, maxiter) + X = rmul!(AP * Y, Diagonal(S⁻¹)) + X .+= X₀ + ΔA = mul!(ΔA, X, Vᴴ, 1, 1) + ΔA = mul!(ΔA, U, Y', 1, 1) end - Yₖᴴ = lmul!(Diagonal(S⁻¹), Xₖ' * AP) - Yₖᴴ .+= Y₀ᴴ - m > n && ((Xₖ, Yₖᴴ) = (Yₖᴴ', Xₖ')) - ΔA = mul!(ΔA, Xₖ, Vᴴ, 1, 1) - ΔA = mul!(ΔA, U, Yₖᴴ, 1, 1) end return ΔA end From 7e53b95c1b7f73cf47daf13a48c5a79ba05a8786 Mon Sep 17 00:00:00 2001 From: leburgel Date: Fri, 2 Oct 2026 13:17:18 +0200 Subject: [PATCH 5/7] Rename helper to `accelerative_smith_iteration!` and add reference --- src/common/pullbacks.jl | 6 ++++-- src/pullbacks/eigh.jl | 4 ++-- src/pullbacks/svd.jl | 4 ++-- 3 files changed, 8 insertions(+), 6 deletions(-) diff --git a/src/common/pullbacks.jl b/src/common/pullbacks.jl index 60b254eca..8257747c4 100644 --- a/src/common/pullbacks.jl +++ b/src/common/pullbacks.jl @@ -37,7 +37,7 @@ is_leading_index(ind::AbstractRange, p::Int) = ind == 1:p is_leading_index(ind::AbstractVector, p::Int) = length(ind) == p && all(ind .== 1:p) """ - _smith_iteration!(X, Xₙ, G, w, atol, maxiter) + accelerative_smith_iteration!(X, Xₙ, G, w, atol, maxiter) Solve `X = B + G * X * Diagonal(w)` by summing the Neumann series `X = Σₖ Gᵏ * B * Diagonal(w)ᵏ` by doubling (Smith's method), i.e. by repeatedly adding @@ -47,8 +47,10 @@ for at most `maxiter` steps. On entry, `X` contains `B`, and it is overwritten with the result. `Xₙ` is used as a buffer, and `G` and `w` are overwritten. It is assumed that `w` is normalized such that `maximum(abs, w) == 1`, so that squaring it can only shrink it. + +Reference: https://doi.org/10.1016/j.aml.2009.01.012. """ -function _smith_iteration!(X, Xₙ, G, w, atol, maxiter) +function accelerative_smith_iteration!(X, Xₙ, G, w, atol, maxiter) Gₙ = similar(G) for k in 1:maxiter Xₙ = rmul!(mul!(Xₙ, G, X), Diagonal(w)) diff --git a/src/pullbacks/eigh.jl b/src/pullbacks/eigh.jl index eaafe820f..9354128e8 100755 --- a/src/pullbacks/eigh.jl +++ b/src/pullbacks/eigh.jl @@ -158,13 +158,13 @@ function eigh_trunc_pullback!( dabsmin = minimum(abs, D) AP ./= dabsmin D⁻¹ = dabsmin ./ D - X = _smith_iteration!(X₀, similar(X₀), AP, D⁻¹, degeneracy_atol, maxiter) + X = accelerative_smith_iteration!(X₀, similar(X₀), AP, D⁻¹, degeneracy_atol, maxiter) Z .+= X # we cannot directly multiply Z * V' into ΔA, because we have to # take the Hermitian part, and cannot apply project_hermitian! to # the current contents of ΔA # TODO: add an `add_project_hermitian!` - # recycle AP's storage, but overwrite it: `_smith_iteration!` may leave a power of AP in it + # recycle AP's storage, but overwrite it: `accelerative_smith_iteration!` may leave a power of AP in it ΔA′ = project_hermitian!(mul!(AP, Z, V')) ΔA .+= ΔA′ else diff --git a/src/pullbacks/svd.jl b/src/pullbacks/svd.jl index d0883c65e..b092782f1 100755 --- a/src/pullbacks/svd.jl +++ b/src/pullbacks/svd.jl @@ -255,7 +255,7 @@ function svd_trunc_pullback!( if m ≤ n X = rmul!(AP * Y₀ᴴ', Diagonal(S⁻¹)) X .+= X₀ - X = _smith_iteration!(X, X₀, AP * AP', S⁻¹ .^ 2, degeneracy_atol, maxiter) # recycle X₀ + X = accelerative_smith_iteration!(X, X₀, AP * AP', S⁻¹ .^ 2, degeneracy_atol, maxiter) # recycle X₀ Yᴴ = lmul!(Diagonal(S⁻¹), X' * AP) Yᴴ .+= Y₀ᴴ ΔA = mul!(ΔA, X, Vᴴ, 1, 1) @@ -263,7 +263,7 @@ function svd_trunc_pullback!( else Y = rmul!(AP' * X₀, Diagonal(S⁻¹)) Y .+= Y₀ᴴ' - Y = _smith_iteration!(Y, similar(Y), AP' * AP, S⁻¹ .^ 2, degeneracy_atol, maxiter) + Y = accelerative_smith_iteration!(Y, similar(Y), AP' * AP, S⁻¹ .^ 2, degeneracy_atol, maxiter) X = rmul!(AP * Y, Diagonal(S⁻¹)) X .+= X₀ ΔA = mul!(ΔA, X, Vᴴ, 1, 1) From b56261898136ab047c824dced1a4ca420c095037 Mon Sep 17 00:00:00 2001 From: Lander Burgelman <39218680+leburgel@users.noreply.github.com> Date: Fri, 2 Oct 2026 15:51:47 +0200 Subject: [PATCH 6/7] Update src/common/pullbacks.jl Co-authored-by: Jutho --- src/common/pullbacks.jl | 3 +++ 1 file changed, 3 insertions(+) diff --git a/src/common/pullbacks.jl b/src/common/pullbacks.jl index 8257747c4..5f5fd6387 100644 --- a/src/common/pullbacks.jl +++ b/src/common/pullbacks.jl @@ -52,6 +52,9 @@ Reference: https://doi.org/10.1016/j.aml.2009.01.012. """ function accelerative_smith_iteration!(X, Xₙ, G, w, atol, maxiter) Gₙ = similar(G) + wmax = maximum(abs, w) + w ./= wmax + G .*= wmax for k in 1:maxiter Xₙ = rmul!(mul!(Xₙ, G, X), Diagonal(w)) if maximum(abs, Xₙ) < atol From ba2e7e72e75a02e565a14a1193715e81b74d988a Mon Sep 17 00:00:00 2001 From: leburgel Date: Fri, 2 Oct 2026 16:12:21 +0200 Subject: [PATCH 7/7] Move normalization inside doubling iteration --- src/common/pullbacks.jl | 4 ++-- src/pullbacks/eigh.jl | 7 +------ src/pullbacks/svd.jl | 4 +--- 3 files changed, 4 insertions(+), 11 deletions(-) diff --git a/src/common/pullbacks.jl b/src/common/pullbacks.jl index 5f5fd6387..cf17a5851 100644 --- a/src/common/pullbacks.jl +++ b/src/common/pullbacks.jl @@ -45,8 +45,8 @@ Solve `X = B + G * X * Diagonal(w)` by summing the Neumann series for at most `maxiter` steps. On entry, `X` contains `B`, and it is overwritten with the result. `Xₙ` is used as a buffer, -and `G` and `w` are overwritten. It is assumed that `w` is normalized such that -`maximum(abs, w) == 1`, so that squaring it can only shrink it. +and `G` and `w` are overwritten. `w` is normalized such that `maximum(abs, w) == 1`, so that +squaring it can only shrink it; `G` is scaled by the inverse factor to compensate. Reference: https://doi.org/10.1016/j.aml.2009.01.012. """ diff --git a/src/pullbacks/eigh.jl b/src/pullbacks/eigh.jl index 9354128e8..12de1452f 100755 --- a/src/pullbacks/eigh.jl +++ b/src/pullbacks/eigh.jl @@ -153,12 +153,7 @@ function eigh_trunc_pullback!( if !iszerotangent(ΔV₊) X₀ = rdiv!(ΔV₊, Diagonal(D)) AP = mul!(copy(A), V * Dmat, V', -1, 1) - # 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 = accelerative_smith_iteration!(X₀, similar(X₀), AP, D⁻¹, degeneracy_atol, maxiter) + X = accelerative_smith_iteration!(X₀, similar(X₀), AP, inv.(D), degeneracy_atol, maxiter) Z .+= X # we cannot directly multiply Z * V' into ΔA, because we have to # take the Hermitian part, and cannot apply project_hermitian! to diff --git a/src/pullbacks/svd.jl b/src/pullbacks/svd.jl index b092782f1..205379aa8 100755 --- a/src/pullbacks/svd.jl +++ b/src/pullbacks/svd.jl @@ -247,9 +247,7 @@ function svd_trunc_pullback!( Y₀ᴴ = iszerotangent(ΔV₊ᴴ) ? zero(Vᴴ) : ldiv!(Diagonal(S), ΔV₊ᴴ) US = mul!(ΔAV, U, Smat) # recycle ΔAV AP = mul!(copy(A), US, Vᴴ, -1, 1) - minS = @view S[end:end] - AP ./= minS - S⁻¹ = minS ./ S + S⁻¹ = inv.(S) # sum the series on the smaller side only, the other side follows from # Yᴴ = Y₀ᴴ + S⁻¹ X' AP (m ≤ n) or X = X₀ + AP Y S⁻¹ (m > n) if m ≤ n