Conversation
Codecov Report✅ All modified and coverable lines are covered by tests.
🚀 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 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 |
…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.
64e4a1d to
bad2b2c
Compare
|
Ok, I did indeed not consider that |
When the cotangents of
svd_pullback!are given on only k 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 k 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
Konly,U₁ᴴΔU₁andV₁ᴴΔV₁are nonzero only in the columnsK. ThenUᴴΔAVis nonzero only in the rows and columnsK, 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 × k block of columnsKofUᴴΔAVand 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 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² k) instead of O(n³). The gauge check covers the same entries as before, since the columnsKcontain 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:k. 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 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: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 k columns (for eigh, the k eigenvalues of largest magnitude) is:ind = Colon())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
mainin 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. Withind = Colon()the result is identical to that ofmain.Benchmark
On
main:With this PR: