Repository navigation
Make svd_pullback! and eigh_pullback! cost proportional to the number of cotangent columns - #292
Conversation
Codecov Report❌ Patch coverage is
... and 5 files with indirect coverage changes 🚀 New features to boost your workflow:
|
80d4af2 to
b0a90bd
Compare
b0a90bd to
fc18c56
Compare
|
The views of They were not the best idea on CPU either, because a
The copies only cost O(n p) next to the O(n² p) 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, p, rng) = (("range", 1:p), ("vector", collect(1:p)), ("perm", randperm(rng, n)[1:p]))
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), p in (50, 200)
A = randn(rng, T, n, p); V = randn(rng, T, n, n); C = zeros(T, n, n)
for (name, ind) in indkinds(n, p, rng)
tv, tc = mintimes(() -> mul!(C, A, view(V, :, ind)', 1, 1), () -> mul!(C, A, V[:, ind]', 1, 1))
@printf("MUL T=%s p=%d ind=%s view=%.3e copy=%.3e view/copy=%.1f\n", T, p, 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 p in (50, 200, 700), (name, ind) in indkinds(n, p, rng)
ΔU = noimagdiag!(randn(rng, T, m, p), U[:, ind])
ΔVᴴ = copy(noimagdiag!(randn(rng, T, n, p), Vᴴ[ind, :]')')
ΔS = Diagonal(randn(rng, real(T), p))
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 p=%d ind=%s views=%.3e copies=%.3e views/copies=%.2f diff=%.0e\n", T, p, name, tv, tc, tv / tc, err)
ΔV = noimagdiag!(randn(rng, T, n, p), V[:, ind]); ΔD = Diagonal(randn(rng, real(T), p))
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 p=%d ind=%s views=%.3e copies=%.3e views/copies=%.2f diff=%.0e\n", T, p, name, tv, tc, tv / tc, err)
end
end
end |
64e4a1d to
bad2b2c
Compare
|
Ok, I did indeed not consider that |
|
Ok, I've pushed some changes. @leburgel, time for you or 🤖 to check whether this is still correct (and of course also CI). |
|
I pushed ba679e9 on top of your typo fixes. Apart from the unpacking of
I also rewrote the comment on the two blocks in 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 |
|
Thanks. The I'll look into the bitvector indexing next, because it is clearly causing issues on the GPU / buildkite CI. |
|
I pushed 2e540de for the remaining Buildkite failure ( |
2e540de to
2e4ec15
Compare
…mber of cotangent columns
Co-authored-by: Jutho <Jutho@users.noreply.github.com>
Co-authored-by: Jutho <Jutho@users.noreply.github.com>
Co-authored-by: Jutho <Jutho@users.noreply.github.com>
Co-authored-by: Jutho <Jutho@users.noreply.github.com>
Co-authored-by: Jutho <Jutho@users.noreply.github.com>
Views with a vector index fall back to scalar indexing in mul! on GPU.
…d_cotangents Co-authored-by: Jutho <Jutho@users.noreply.github.com>
2e4ec15 to
8051d01
Compare
|
Buildkite shows a failure in the |

When the cotangents of
svd_pullback!are given on only p of the r singular vectors (throughind), the pullback pads them with zeros to all r columns and continues with r × r matrices:check_and_prepare_svd_cotangentsformsU₁' * ΔU₁andΔU₁ - U₁ * (U₁' * ΔU₁), the same forV, and the result is applied asU₁ * (UᴴΔAV * V₁ᴴ). The cost is therefore O(m n r) whatever p is.eigh_pullback!does the same withV' * ΔV₁andV * 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, whereindholds the kept indices.With cotangents on the columns
indonly,U₁ᴴΔU₁andV₁ᴴΔV₁are nonzero only in the columnsind. ThenUᴴΔAVis nonzero only in the rows and columnsind, and its rows follow from its columns by antihermiticity. In this PR,check_and_prepare_svd_cotangentstherefore no longer pads the cotangents, but computes only the r × p block of columnsindofUᴴΔAVand the corresponding block of its rows. For 2p ≤ r,svd_pullback!applies these two blocks directly as rank-p updates, at cost O(m n p); otherwise it assemblesUᴴΔAVfrom them and applies it as before.check_and_prepare_eigh_cotangentsandeigh_pullback!are changed in the same way, so that for 2k ≤ n the cost ofeigh_pullback!is O(n² p) instead of O(n³). The gauge check covers the same entries as before, since the columnsindcontain every nonzero entry up to conjugation. Cotangents on columns beyond a full rank have components along all ofU₁orV₁ᴴ, so in that case the block still spans all r columns.svd_trunc_pullback!andeigh_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ΔSwith anindother than1:p. Since #232,check_and_prepare_svd_cotangentsonmainindexesΔSby column number instead of by position withinind, so that for exampleind = [3, 1, 7, 2]throws aBoundsError. The new code adds each entry ofΔSto 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]and1:6, and with zeroΔU,ΔVᴴorΔS, the result agrees withmain(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 p between n / 36 and n / 9 (the spectra and sizes of a CTMRG step in PEPSKit.jl, with n = χD² and p = χ, which motivated this investigation), it is 4–22× faster:main(s)Minimum times on a laptop with 4 BLAS threads.
main'ssrc/pullbacks/svd.jlandsrc/pullbacks/eigh.jlwere 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
mainfor cotangents on the first p columns (for eigh, the p eigenvalues of largest magnitude) is:ind = Colon())With 2p ≤ r the new path is faster in every case. Above that, the r × r matrix is formed from the p computed columns and applied as before, which is slower than
mainin five of the 24 cases with 2p > r: by up to 1.3× in two SVD cases at n = 1000, and by at most 6% in the other three. Withind = Colon()the result is identical to that ofmain.Benchmark
On
main:With this PR: