Skip to content

Make svd_pullback! and eigh_pullback! cost proportional to the number of cotangent columns - #292

Open
leburgel wants to merge 21 commits into
mainfrom
lb/full_pullback_kept_columns
Open

leburgel wants to merge 21 commits into
mainfrom
lb/full_pullback_kept_columns

Conversation

@leburgel

@leburgel leburgel commented Oct 1, 2026

Copy link
Copy Markdown
Member

When the cotangents of svd_pullback! are given on only k of the r singular vectors (through ind), the pullback pads them with zeros to all r columns and continues with r × r matrices: check_and_prepare_svd_cotangents forms U₁' * ΔU₁ and ΔU₁ - U₁ * (U₁' * ΔU₁), the same for V, and the result is applied as U₁ * (UᴴΔAV * V₁ᴴ). The cost is therefore O(m n r) whatever k is. eigh_pullback! does the same with V' * ΔV₁ and V * VᴴΔAV * V', at cost O(n³) for an n × n matrix. This is the common case for a full pullback of a truncated decomposition, where ind holds the kept indices.

With cotangents on the columns K only, U₁ᴴΔU₁ and V₁ᴴΔV₁ are nonzero only in the columns K. Then UᴴΔAV is nonzero only in the rows and columns K, and its rows follow from its columns by antihermiticity. In this PR, check_and_prepare_svd_cotangents therefore no longer pads the cotangents, but computes only the r × k block of columns K of UᴴΔAV and the corresponding block of its rows. For 2k ≤ r, svd_pullback! applies these two blocks directly as rank-k updates, at cost O(m n k); otherwise it assembles UᴴΔAV from them and applies it as before. check_and_prepare_eigh_cotangents and eigh_pullback! are changed in the same way, so that for 2k ≤ n the cost of eigh_pullback! is O(n² k) instead of O(n³). The gauge check covers the same entries as before, since the columns K contain every nonzero entry up to conjugation. Cotangents on columns beyond a full rank have components along all of U₁ or V₁ᴴ, so in that case the block still spans all r columns. svd_trunc_pullback! and eigh_trunc_pullback! call the same functions with all columns and receive the same full matrices as before.

Bug fix. This PR also fixes svd_pullback! for a nonzero ΔS with an ind other than 1:k. Since #232, check_and_prepare_svd_cotangents on main indexes ΔS by column number instead of by position within ind, so that for example ind = [3, 1, 7, 2] throws a BoundsError. The new code adds each entry of ΔS to the diagonal entry of its own column; the comparison below includes this case.

On random real and complex, square and rectangular, full-rank and rank-deficient matrices, with ind = 1:5, [3, 1, 7, 2] and 1:6, and with zero ΔU, ΔVᴴ or ΔS, the result agrees with main (called with the same cotangents zero-padded to all columns) to 7.9e-16 relative. On square matrices with exponentially decaying singular values or eigenvalues and k between n / 36 and n / 9 (the spectra and sizes of a CTMRG step in PEPSKit.jl, with n = χD² and k = χ, which motivated this investigation), it is 4–22× faster:

decomposition eltype n k main (s) this PR (s) speedup
SVD Float64 180 20 1.30e-3 3.29e-4 3.9×
SVD Float64 640 40 4.57e-2 5.51e-3 8.3×
SVD Float64 1500 60 0.515 2.34e-2 22.0×
SVD Float64 2880 80 2.31 0.173 13.3×
SVD ComplexF64 180 20 4.40e-3 8.23e-4 5.3×
SVD ComplexF64 640 40 0.175 1.67e-2 10.5×
SVD ComplexF64 1500 60 1.10 7.34e-2 14.9×
SVD ComplexF64 2880 80 9.12 0.622 14.7×
eigh Float64 180 20 8.06e-4 1.97e-4 4.1×
eigh Float64 640 40 2.30e-2 3.23e-3 7.1×
eigh Float64 1500 60 0.276 2.34e-2 11.8×
eigh Float64 2880 80 1.58 0.121 13.1×
eigh ComplexF64 180 20 2.27e-3 5.30e-4 4.3×
eigh ComplexF64 640 40 8.68e-2 1.04e-2 8.3×
eigh ComplexF64 1500 60 0.894 8.92e-2 10.0×
eigh ComplexF64 2880 80 4.72 0.240 19.6×

Minimum times on a laptop with 4 BLAS threads. main's src/pullbacks/svd.jl and src/pullbacks/eigh.jl were loaded into the same process, and the two versions were timed alternately with the same arguments.

On random n × n matrices (n = 400 and 1000, real and complex; same timing method), the speedup over main for cotangents on the first k columns (for eigh, the k eigenvalues of largest magnitude) is:

k / n path SVD eigh
1/20 new 9.3–14× 7.5–16×
1/4 new 2.7–3.3× 2.3–2.4×
9/20 new 1.55–1.67× 1.30–1.76×
1/2 new 1.37–1.50× 1.19–1.23×
11/20 dense 1.34–1.53× 1.15–1.18×
7/10 dense 0.76–1.21× 0.94–1.25×
9/10 dense 0.77–1.10× 0.97–1.04×
1 (ind = Colon()) dense 0.99–1.14× 1.00–1.16×

With 2k ≤ r the new path is faster in every case. Above that, the r × r matrix is formed from the k computed columns and applied as before, which is slower than main in five of the 24 cases with 2k > r: by up to 1.3× in two SVD cases at n = 1000, and by at most 6% in the other three. With ind = Colon() the result is identical to that of main.

Benchmark
using LinearAlgebra, Random, Printf
using MatrixAlgebraKit
using MatrixAlgebraKit: svd_pullback!, eigh_pullback!, diagview

# gauge-invariant cotangents on the columns `ind` (the imaginary part of diag(Xᴴ ΔX) is the gauge part)
noimagdiag!(ΔX, X) = (ΔX .-= X .* transpose(im .* imag.(diag(X' * ΔX))); ΔX)
noimagdiag!(ΔX::AbstractMatrix{<:Real}, X) = ΔX
pad(X, ind, p) = (Y = zeros(eltype(X), size(X, 1), p); Y[:, ind] .= X; Y)
padvec(x, ind, p) = (y = zeros(eltype(x), p); y[ind] .= x; y)

# time of the pullback with cotangents on the columns ind, and its difference with the same
# cotangents zero-padded to all columns
function timeit(name, pullback!, A, X, cots, padded, ind)
    f = () -> pullback!(zero(A), A, X, cots, ind)
    t = (f(); minimum(@elapsed(f()) for _ in 1:5))
    ref = pullback!(zero(A), A, X, padded)
    @printf("%-4s %-10s n = %4d  k = %2d  time %.2e s  difference %.1e\n",
        name, eltype(A), size(A, 1), length(ind), t, norm(f() - ref) / norm(ref))
end

BLAS.set_num_threads(4)
rng = Xoshiro(1)
for T in (Float64, ComplexF64), (n, k) in ((180, 20), (640, 40), (1500, 60), (2880, 80))
    ind = 1:k
    Q₁, Q₂ = Matrix(qr(randn(rng, T, n, n)).Q), Matrix(qr(randn(rng, T, n, n)).Q)

    A = Q₁ * Diagonal(max.(0.98 .^ (0:(n - 1)), 1.0e-10)) * Q₂'
    U, S, Vᴴ = svd_compact(A)
    ΔU = noimagdiag!(randn(rng, T, n, k), U[:, ind])
    ΔV = noimagdiag!(randn(rng, T, n, k), Vᴴ[ind, :]')
    ΔS = randn(rng, real(T), k)
    timeit("svd", svd_pullback!, A, (U, S, Vᴴ), (ΔU, Diagonal(ΔS), copy(ΔV')),
        (pad(ΔU, ind, n), Diagonal(padvec(ΔS, ind, n)), copy(pad(ΔV, ind, n)')), ind)

    λ = max.(0.956 .^ (0:(n - 1)), 1.0e-10) .* rand(rng, (-1, 1), n)
    H = Matrix(Hermitian(Q₁ * Diagonal(λ) * Q₁'))
    D, V = eigh_full(H)
    indh = sort(sortperm(abs.(diagview(D)); rev = true)[1:k])
    ΔV = noimagdiag!(randn(rng, T, n, k), V[:, indh])
    ΔD = randn(rng, real(T), k)
    timeit("eigh", eigh_pullback!, H, (D, V), (Diagonal(ΔD), ΔV),
        (Diagonal(padvec(ΔD, indh, n)), pad(ΔV, indh, n)), indh)
end

On main:

svd  Float64    n =  180  k = 20  time 1.96e-03 s  difference 0.0e+00
eigh Float64    n =  180  k = 20  time 5.86e-04 s  difference 0.0e+00
svd  Float64    n =  640  k = 40  time 2.66e-02 s  difference 0.0e+00
eigh Float64    n =  640  k = 40  time 1.22e-02 s  difference 0.0e+00
svd  Float64    n = 1500  k = 60  time 3.78e-01 s  difference 0.0e+00
eigh Float64    n = 1500  k = 60  time 2.88e-01 s  difference 0.0e+00
svd  Float64    n = 2880  k = 80  time 3.08e+00 s  difference 0.0e+00
eigh Float64    n = 2880  k = 80  time 1.27e+00 s  difference 0.0e+00
svd  ComplexF64 n =  180  k = 20  time 2.18e-03 s  difference 0.0e+00
eigh ComplexF64 n =  180  k = 20  time 2.29e-03 s  difference 0.0e+00
svd  ComplexF64 n =  640  k = 40  time 1.79e-01 s  difference 0.0e+00
eigh ComplexF64 n =  640  k = 40  time 8.61e-02 s  difference 0.0e+00
svd  ComplexF64 n = 1500  k = 60  time 1.83e+00 s  difference 0.0e+00
eigh ComplexF64 n = 1500  k = 60  time 5.79e-01 s  difference 0.0e+00
svd  ComplexF64 n = 2880  k = 80  time 9.11e+00 s  difference 0.0e+00
eigh ComplexF64 n = 2880  k = 80  time 3.91e+00 s  difference 0.0e+00

With this PR:

svd  Float64    n =  180  k = 20  time 3.24e-04 s  difference 1.1e-15
eigh Float64    n =  180  k = 20  time 1.98e-04 s  difference 7.4e-16
svd  Float64    n =  640  k = 40  time 5.45e-03 s  difference 1.2e-15
eigh Float64    n =  640  k = 40  time 3.45e-03 s  difference 7.2e-16
svd  Float64    n = 1500  k = 60  time 4.00e-02 s  difference 1.1e-15
eigh Float64    n = 1500  k = 60  time 2.41e-02 s  difference 7.1e-16
svd  Float64    n = 2880  k = 80  time 2.17e-01 s  difference 1.0e-15
eigh Float64    n = 2880  k = 80  time 8.40e-02 s  difference 7.7e-16
svd  ComplexF64 n =  180  k = 20  time 7.32e-04 s  difference 1.4e-15
eigh ComplexF64 n =  180  k = 20  time 5.02e-04 s  difference 7.8e-16
svd  ComplexF64 n =  640  k = 40  time 1.39e-02 s  difference 1.5e-15
eigh ComplexF64 n =  640  k = 40  time 4.62e-03 s  difference 9.3e-16
svd  ComplexF64 n = 1500  k = 60  time 1.39e-01 s  difference 1.4e-15
eigh ComplexF64 n = 1500  k = 60  time 4.81e-02 s  difference 9.9e-16
svd  ComplexF64 n = 2880  k = 80  time 6.47e-01 s  difference 1.2e-15
eigh ComplexF64 n = 2880  k = 80  time 2.13e-01 s  difference 9.4e-16

@leburgel
leburgel requested a review from Jutho October 1, 2026 11:53
Comment thread src/common/pullbacks.jl
@codecov

codecov Bot commented Oct 1, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.

Files with missing lines Coverage Δ
src/common/pullbacks.jl 94.11% <100.00%> (+1.26%) ⬆️
src/pullbacks/eigh.jl 87.20% <100.00%> (+1.13%) ⬆️
src/pullbacks/svd.jl 95.43% <100.00%> (+0.95%) ⬆️
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

Comment thread src/pullbacks/eigh.jl Outdated
Comment thread src/pullbacks/eigh.jl Outdated
Comment thread src/pullbacks/eigh.jl Outdated
Comment thread src/pullbacks/eigh.jl Outdated
@leburgel
leburgel force-pushed the lb/full_pullback_kept_columns branch from 80d4af2 to b0a90bd Compare October 2, 2026 14:21
Comment thread src/pullbacks/eigh.jl Outdated
Comment thread src/pullbacks/svd.jl Outdated
@leburgel
leburgel force-pushed the lb/full_pullback_kept_columns branch from b0a90bd to fc18c56 Compare October 3, 2026 07:32
@leburgel

leburgel commented Oct 3, 2026

Copy link
Copy Markdown
Member Author

The views of U, Vᴴ and V broke the GPU tests (the failing Buildkite mooncake jobs): when ind′ is a Vector{Int}, the view is not a StridedMatrix, so mul! falls back to the generic matrix product, which is scalar indexing on a CuArray. I therefore went back to copies, renamed to Uₖ, Vᴴₖ and Vₖ like Dₖ, and kept the view of S, which is only broadcast.

They were not the best idea on CPU either, because a Vector ind′ is common: truncerror, and truncrank for eigh, return one even when the kept columns are contiguous. The ratio of times views/copies for the pullbacks (laptop, 4 BLAS threads, Float64 and ComplexF64, svd of 1200×1000, eigh of 1000×1000, k = 50, 200, 700):

ind a range ind a Vector
svd_pullback! 0.98–1.18 1.7–7.0
eigh_pullback! 0.99–1.01 3.3–9.5 (1.0 for k > n/2, which does not use Vₖ)

The copies only cost O(n k) next to the O(n² k) products.

Benchmark script (run on fc18c56, the version with views)
# Views vs copies of the kept columns in svd_pullback!/eigh_pullback! (untracked, for the review reply):
#   1. mul! with a view of the columns ind as one factor, for ind a range and a Vector
#   2. the pullbacks of this branch (views of U, Vᴴ and V) against the same code with copies
#      `U[:, ind′]`, `Vᴴ[ind′, :]`, `V[:, ind′]` (module Copies), timed alternately in one process
using LinearAlgebra, Random, Printf
using MatrixAlgebraKit
using MatrixAlgebraKit: svd_pullback!, eigh_pullback!
BLAS.set_num_threads(4)
module Copies
using LinearAlgebra, MatrixAlgebraKit
const SUBS = ("view(U, :, ind′)" => "U[:, ind′]", "view(Vᴴ, ind′, :)" => "Vᴴ[ind′, :]", "view(V, :, ind′)" => "V[:, ind′]")
const SRC = map(("svd", "eigh")) do f
    src = read(joinpath(pkgdir(MatrixAlgebraKit), "src", "pullbacks", "$f.jl"), String)
    any(occursin(first(s), src) for s in SUBS) || error("no view in $f.jl")
    return replace(src, SUBS...)
end
const OWN = Set(Symbol(something(m[1], m[2])) for src in SRC for m in eachmatch(r"^\s*function ([^\s(]+)\(|^(?:function )?([A-Za-z_][^\s(]*)\([^\n]*\)\s*=[^=]"m, src))
for name in names(MatrixAlgebraKit; all = true)
    (name in OWN || startswith(string(name), "#") || isdefined(@__MODULE__, name)) && continue
    @eval const $name = MatrixAlgebraKit.$name
end
foreach(src -> include_string(@__MODULE__, src), SRC)
end
function mintimes(f, g; mintotal = 2.0, minreps = 5, maxreps = 15)
    f(); g(); tf = tg = Inf; tot = 0.0
    for i in 1:maxreps
        s = time_ns(); f(); df = (time_ns() - s) / 1.0e9
        s = time_ns(); g(); dg = (time_ns() - s) / 1.0e9
        tf, tg = min(tf, df), min(tg, dg); tot += df + dg
        i >= minreps && tot > mintotal && break
    end
    return tf, tg
end
noimagdiag!(ΔX, X) = (ΔX .-= X .* transpose(im .* imag.(diag(X' * ΔX))); ΔX)
noimagdiag!(ΔX::AbstractMatrix{<:Real}, X) = ΔX
indkinds(n, k, rng) = (("range", 1:k), ("vector", collect(1:k)), ("perm", randperm(rng, n)[1:k]))

println("== mul!(C, A, B') with B = view(V, :, ind) or V[:, ind], n = 1000, BLAS threads 4")
let rng = Xoshiro(1), n = 1000
    for T in (Float64, ComplexF64), k in (50, 200)
        A = randn(rng, T, n, k); V = randn(rng, T, n, n); C = zeros(T, n, n)
        for (name, ind) in indkinds(n, k, rng)
            tv, tc = mintimes(() -> mul!(C, A, view(V, :, ind)', 1, 1), () -> mul!(C, A, V[:, ind]', 1, 1))
            @printf("MUL T=%s k=%d ind=%s view=%.3e copy=%.3e view/copy=%.1f\n", T, k, name, tv, tc, tv / tc)
        end
    end
end

println("== pullbacks: views (branch) vs copies, svd m=1200 n=1000, eigh n=1000")
let rng = Xoshiro(2), m = 1200, n = 1000
    for T in (Float64, ComplexF64)
        A = randn(rng, T, m, n); U, S, Vᴴ = svd_compact(A); Sd = diag(S)
        H = randn(rng, T, n, n); H = H + H'; D, V = eigh_full(H); Dd = diag(D)
        for k in (50, 200, 700), (name, ind) in indkinds(n, k, rng)
            ΔU = noimagdiag!(randn(rng, T, m, k), U[:, ind])
            ΔVᴴ = copy(noimagdiag!(randn(rng, T, n, k), Vᴴ[ind, :]')')
            ΔS = Diagonal(randn(rng, real(T), k))
            f(pb) = () -> pb(zero(A), A, (U, S, Vᴴ), (ΔU, ΔS, ΔVᴴ), ind)
            err = norm(f(svd_pullback!)() - f(Copies.svd_pullback!)())
            tv, tc = mintimes(f(svd_pullback!), f(Copies.svd_pullback!))
            @printf("PB dec=svd T=%s k=%d ind=%s views=%.3e copies=%.3e views/copies=%.2f diff=%.0e\n", T, k, name, tv, tc, tv / tc, err)
            ΔV = noimagdiag!(randn(rng, T, n, k), V[:, ind]); ΔD = Diagonal(randn(rng, real(T), k))
            g(pb) = () -> pb(zero(H), H, (D, V), (ΔD, ΔV), ind)
            err = norm(g(eigh_pullback!)() - g(Copies.eigh_pullback!)())
            tv, tc = mintimes(g(eigh_pullback!), g(Copies.eigh_pullback!))
            @printf("PB dec=eigh T=%s k=%d ind=%s views=%.3e copies=%.3e views/copies=%.2f diff=%.0e\n", T, k, name, tv, tc, tv / tc, err)
        end
    end
end

@leburgel
leburgel force-pushed the lb/full_pullback_kept_columns branch from 64e4a1d to bad2b2c Compare October 4, 2026 07:42
@Jutho

Jutho commented Oct 4, 2026

Copy link
Copy Markdown
Member

Ok, I did indeed not consider that ind is a Vector{Int} so often.

Comment thread src/pullbacks/svd.jl Outdated
Comment on lines +24 to +26
fold = r == minmn && max(length(indU), length(indV)) > length(j₁)
ind′ = fold ? (1:r) : indS[j₁]
l₁ = fold ? indS[j₁] : eachindex(j₁) # columns of the cotangents j₁ among ind′

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 am very confused by what this is doing. I understand the first line, j₁ is such that all of indS[j₁] is in 1:r. But I am confused by what fold means. I am even more confused by l₁, which is either indS[j₁] or eachindex(j₁) which seem to very different things. In the case where j₁ is eachindex(S), it means that l₁ is then either indS[eachindex(indS)], which is simply indS, or alternatively eachindex(eachindex(indS)), of which I have idea what this means.

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.

You're right, this was hard to follow. The cotangents are given for the columns indS, and j₁ are the positions in indS of those within the rank. From that:

  • ind′ are the columns of UᴴΔAV that are computed, normally the cotangent columns within the rank, indS[j₁].
  • l₁ are the positions of the cotangents j₁ among those columns, so that ind′[l₁] == indS[j₁]; they place ΔU[:, j₁] in ΔU₁ and ΔS[j₁] on the diagonal. With ind′ = indS[j₁] these are just 1:length(j₁), i.e. eachindex(j₁), which would be clearer as eachindex(ind′).

fold is the exception, the r == minmn branch of the loop on main: with full rank and cotangents beyond it, their components along U₁ are folded into every column of ΔU₁, so ind′ = 1:r and the positions are the indices themselves, l₁ = indS[j₁]. As far as I can see, this only happens for svd_full output with ind = Colon() and m ≠ n, since a vector ind beyond min(m, n) already errors at axes(S, 1)[ind].

So for j₁ = eachindex(indS): without fold, ind′ = indS and l₁ = 1:length(indS); with fold, ind′ = 1:r and l₁ = indS.

I've rewritten the comment along these lines in 7b2d472.

Comment thread src/pullbacks/svd.jl Outdated
Comment thread src/pullbacks/svd.jl Outdated
else
ΔU₁ = zero(U₁)
ΔU₁ = zero!(similar(U, (m, k)))
ΔU₁[:, l₁] .= view(ΔU, :, j₁)

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 will try to rewrite this myself, because currently it is very confusing and the naming is wrong. The subscript 1 refers to the 1:r range in e.g. U₁. In the eigh code, a subscript k was used for some of the quantities that had columns associated with ind′ of length k, though also not consistently. But trying to make that consistent, this should be called ΔUₖ. But I think a larger overhaul of the structure is necessary.

@Jutho

Jutho commented Oct 5, 2026

Copy link
Copy Markdown
Member

Ok, I've pushed some changes. @leburgel, time for you or 🤖 to check whether this is still correct (and of course also CI).

@leburgel

leburgel commented Oct 5, 2026

Copy link
Copy Markdown
Member Author

I pushed ba679e9 on top of your typo fixes. Apart from the unpacking of ΔDV in eigh_pullback! (it now takes ind₀ from check_and_prepare_eigh_cotangents) and VᴴΔV₀' → VᴴΔV₁₀', two things gave wrong results:

  • check_and_prepare_svd_cotangents returned aUᴴΔU₁₀ instead of aUᴴΔAV₁₀, the antihermitian part of UᴴΔAV; the call sites now use that name.
  • The columns ind₀ of UᴴΔAV are hUᴴΔAV₁₀ .+ aUᴴΔAV₁₀, without adjoints; only the rows are hUᴴΔAV₁₀' .- aUᴴΔAV₁₀'.

I also rewrote the comment on the two blocks in svd_pullback!, which still described UᴴΔAVₖ and UᴴΔAVʳ; feel free to reword it.

Two changes in behaviour compared with main, which I left as they are: cotangents beyond the rank in the rank-deficient case are no longer checked to be zero, and for svd_full with ind = Colon() and rank below min(m, n) they now throw the ArgumentError. The tests don't cover the latter. Also, ΔU[:, J] and view(ΔS, J) index with a BitVector, which may need findall(J) on GPU.

@Jutho

Jutho commented Oct 5, 2026

Copy link
Copy Markdown
Member

Thanks. The aUᴴΔU -> aUᴴΔAV bug makes me wonder if this was something that VSCode autocorrect or copilot changed behind my back.

I'll look into the bitvector indexing next, because it is clearly causing issues on the GPU / buildkite CI.

@leburgel

leburgel commented Oct 6, 2026

Copy link
Copy Markdown
Member Author

I pushed 2e540de for the remaining Buildkite failure (svd_full in cuda / mooncake). For svd_full, ΔS = diagview(ΔSmat) is itself a view of the dense ΔSmat, so view(ΔS, J₁) on line 93 was a view of a view with a CPU Vector{Int} index, which cannot be passed to the GPU kernel; for svd_compact, ΔS is the vector of a Diagonal, which is why only svd_full failed. It is now ΔS[J₁], as S[ind₀] for S₀. I did the same for view(ΔS, J₂) on line 94 and for S₀ = view(S, ind₀) in svd_pullback!, which is used in the broadcasts with ΔU₊ and ΔV₊ᴴ, and applied Runic's 0:(length(ind₀) - 1) on line 93.

@leburgel
leburgel force-pushed the lb/full_pullback_kept_columns branch from 2e540de to 2e4ec15 Compare October 6, 2026 15:44
Comment thread src/pullbacks/svd.jl
length(indS) == length(ΔS) || throw(DimensionMismatch(lazy"length of selected S values ($(length(indS))) does not match length of ΔS ($(length(ΔS)))"))
hUᴴΔAV₁₀[ind₀ .+ r .* (0:length(ind₀)-1)] .+= real.(view(ΔS, J₁)) # diagonal entries
Δgauge = max(Δgauge, maximum(abs, view(ΔS, J₂); init = zero(Δgauge)))
hUᴴΔAV₁₀[ind₀ .+ r .* (0:(length(ind₀) - 1))] .+= real.(ΔS[J₁]) # diagonal entries

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.

Is this fixing the GPU. Thanks, this was the next thing I was going to try, but it was very much hit and miss so far. What is strange, is that the original code that you/Claude generated also had a view(ΔS, l₁) on this location, so I hoped this view would have worked.

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.

Also, the corresponding line in the eigh pullback still has a view. I guess for one or the other reason, there the corresponding indices are always a range in each of the test cases, but this will probably break as soon as someone comes up with a truncation strategy that does not produce a range.

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.

Oh I only now read the comment, so the problem is with the view of a view in the case of svd_full. That's why eigh is not causing problems. I have a hard time reading the buildkite log files, and downloading them locally results in all kinds of weird Unicode garbage.

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.

Maybe then it is worth to try to intercept the svd_full case, something like

ΔS₀ = (ΔS isa SubArray) ? view(ΔSmat, J₁ .+ (J₁ .- 1) .* size(ΔSmat, 1)) : view(ΔS, J₁)

Then again, instantiating the broadcast J₁ .+ (J₁ .- 1) .* size(ΔSmat, 1) probably allocates as much memory as slicing ΔS[J₁], so I guess the current solution is fine.

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.

2 participants