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 |
64e4a1d to
bad2b2c
Compare
|
Ok, I did indeed not consider that |
| 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′ |
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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 ofUᴴΔAVthat are computed, normally the cotangent columns within the rank,indS[j₁].l₁are the positions of the cotangentsj₁among those columns, so thatind′[l₁] == indS[j₁]; they placeΔU[:, j₁]inΔU₁andΔS[j₁]on the diagonal. Withind′ = indS[j₁]these are just1:length(j₁), i.e.eachindex(j₁), which would be clearer aseachindex(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.
| else | ||
| ΔU₁ = zero(U₁) | ||
| ΔU₁ = zero!(similar(U, (m, k))) | ||
| ΔU₁[:, l₁] .= view(ΔU, :, j₁) |
There was a problem hiding this comment.
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.
|
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 ( |
…mber of cotangent columns
Co-authored-by: Jutho <[email protected]>
Co-authored-by: Jutho <[email protected]>
Co-authored-by: Jutho <[email protected]>
Co-authored-by: Jutho <[email protected]>
Co-authored-by: Jutho <[email protected]>
Views with a vector index fall back to scalar indexing in mul! on GPU.
…d_cotangents Co-authored-by: Jutho <[email protected]>
2e540de to
2e4ec15
Compare
| 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 |
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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.
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: