Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,9 @@ TensorKitFiniteDifferencesExt = "FiniteDifferences"
TensorKitGPUArraysExt = "GPUArrays"
TensorKitMooncakeExt = "Mooncake"

[sources]
MatrixAlgebraKit = {url = "https://github.com/QuantumKitHub/MatrixAlgebraKit.jl", rev = "main"}

[compat]
AMDGPU = "2"
Adapt = "4"
Expand Down
2 changes: 1 addition & 1 deletion ext/TensorKitGPUArraysExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -411,7 +411,7 @@ end
st_dst, offs_dst = @inbounds structs_dst[blk.dst_offset + i + 1]
i_dst = _linear_index(coords, st_dst, offs_dst)

# dst_i = β * dst_i + α * Σ_j U[i, j] * permute(src_j, p): each output tree is a
# dst_i = β * dst_i + α * Σ_j U[i, j] * permute(op(src_j), p): each output tree is a
# linear combination of the input trees weighted by the recoupling coefficients.
# The permutation of src_j was already done by permuting its strides before the
# kernel launched.
Expand Down
76 changes: 71 additions & 5 deletions src/factorizations/matrixalgebrakit.jl
Original file line number Diff line number Diff line change
@@ -1,5 +1,16 @@
# Algorithm selection
# -------------------

"""
_tensor_algorithm(f!, ::Type{<:AbstractTensorMap}; kwargs...)

Algorithm a `TensorMap` factorization defaults to. Blocks are decomposed one at a time, so
the default is whatever the block type would use. For algorithms that have a batched version,
like `QRIteration` or `Jacobi`, this can be overridden to point to the batching version.
"""
function _tensor_algorithm(f!, ::Type{T}; kwargs...) where {T <: AbstractTensorMap}
return MAK.default_algorithm(f!, blocktype(T); kwargs...)
end
for f in
[
:svd_compact, :svd_full, :svd_vals,
Expand All @@ -12,7 +23,7 @@ for f in
]
f! = Symbol(f, :!)
@eval function MAK.default_algorithm(::typeof($f!), ::Type{T}; kwargs...) where {T <: AbstractTensorMap}
return MAK.default_algorithm($f!, blocktype(T); kwargs...)
return _tensor_algorithm($f!, T; kwargs...)
end
@eval function MAK.copy_input(::typeof($f), t::AbstractTensorMap)
return @timeit_debug GLOBAL_TIMER "alloc: copy_input" copy_oftype(
Expand All @@ -35,8 +46,7 @@ end
# -----------------------
for f! in (
:qr_compact!, :qr_full!, :lq_compact!, :lq_full!,
:eig_full!, :eigh_full!, :svd_compact!, :svd_full!,
:left_polar!, :right_polar!,
:eig_full!, :eigh_full!, :left_polar!, :right_polar!,
)
@eval function MAK.$f!(t::AbstractTensorMap, F, alg::AbstractAlgorithm)
$(f! in (:eig_full!, :eigh_full!) && :(LinearAlgebra.checksquare(t)))
Expand All @@ -56,10 +66,40 @@ for f! in (
end
end

# these have batched versions, use them if the driver supports it
# which blocks are worth batching, and how to pack them, is decided by MatrixAlgebraKit
for (f!, bf!) in ((:svd_compact!, :batched_svd_compact!), (:svd_full!, :batched_svd_full!))
@eval function MAK.$f!(t::AbstractTensorMap, F, alg::AbstractAlgorithm)
@timeit_debug GLOBAL_TIMER $(string(f!)) begin
U, S, Vᴴ = F
driver = get(alg.kwargs, :driver, MAK.DefaultDriver())
if MAK.supports_ragged_batch(MAK.$bf!, alg, driver, storagetype(t))
cs = collect(blocksectors(t))
As = [block(t, c) for c in cs]
Fs = ([block(U, c) for c in cs], [block(S, c) for c in cs], [block(Vᴴ, c) for c in cs])
if applicable(MAK.$bf!, As, Fs, alg)
@timeit_debug GLOBAL_TIMER "batched: MatrixAlgebraKit" MAK.$bf!(As, Fs, alg)
return F
end
end
foreachblock(t, U, S, Vᴴ) do _, (tblock, Fblocks...)
@timeit_debug GLOBAL_TIMER "dense: MatrixAlgebraKit" begin
Fblocks′ = $f!(tblock, Fblocks, alg)
# deal with the case where the output is not in-place
for (b′, b) in zip(Fblocks′, Fblocks)
b === b′ || copy!(b, b′)
end
end
return nothing
end
return F
end
end
end

# Handle these separately because single output instead of tuple
for f! in (
:qr_null!, :lq_null!,
:svd_vals!, :eig_vals!, :eigh_vals!,
:qr_null!, :lq_null!, :eig_vals!, :eigh_vals!,
:project_hermitian!, :project_antihermitian!, :project_isometric!,
:exponential!,
)
Expand All @@ -79,6 +119,32 @@ for f! in (
end
end

# Handle these separately because single output instead of tuple AND batching is supported
for (f!, bf!) in ((:svd_vals!, :batched_svd_vals!),)
@eval function MAK.$f!(t::AbstractTensorMap, N, alg::AbstractAlgorithm)
@timeit_debug GLOBAL_TIMER $(string(f!)) begin
driver = get(alg.kwargs, :driver, MAK.DefaultDriver())
if MAK.supports_ragged_batch(MAK.$bf!, alg, driver, storagetype(t))
cs = collect(blocksectors(t))
As, Ns = [block(t, c) for c in cs], [block(N, c) for c in cs]
if applicable(MAK.$bf!, As, Ns, alg)
@timeit_debug GLOBAL_TIMER "batched: MatrixAlgebraKit" MAK.$bf!(As, Ns, alg)
return N
end
end
foreachblock(t, N) do _, (tblock, Nblock)
@timeit_debug GLOBAL_TIMER "dense: MatrixAlgebraKit" begin
Nblock′ = $f!(tblock, Nblock, alg)
# deal with the case where the output is not the same as the input
Nblock === Nblock′ || copy!(Nblock, Nblock′)
end
return nothing
end
end
return N
end
end

# Exponential with Tuple
function MAK.exponential!((τ, t)::Tuple{E, T}, N, alg::AbstractAlgorithm) where {E <: Number, T <: AbstractTensorMap}
LinearAlgebra.checksquare(t)
Expand Down
Loading