diff --git a/docs/src/Changelog.md b/docs/src/Changelog.md index e3baac971..6c662ca1a 100644 --- a/docs/src/Changelog.md +++ b/docs/src/Changelog.md @@ -24,6 +24,9 @@ When releasing a new version, move the "Unreleased" changes to a new version sec ### Changed +- `@cached` can be used from other modules, such as package extensions, which then own the cache of the methods they add; these caches are shown with their module in `global_cache_info` (internal) +- The `TreeTransformer`s used in index manipulations are now constructed and cached per storagetype of the destination tensor, which is also a dispatch point for storage-specific transformers with their own cache (internal) + ### Deprecated ### Removed @@ -32,6 +35,8 @@ When releasing a new version, move the "Unreleased" changes to a new version sec ### Performance +- Recoupling matrices are stored such that they can leverage BLAS, also for mixed complex tensor with real recoupling coefficients. + ## [0.17.2](https://github.com/QuantumKitHub/TensorKit.jl/compare/v0.17.1...v0.17.2) - 2026-09-20 ### Added diff --git a/src/auxiliary/caches.jl b/src/auxiliary/caches.jl index f3bec7844..e77b96a72 100644 --- a/src/auxiliary/caches.jl +++ b/src/auxiliary/caches.jl @@ -54,9 +54,26 @@ function _cached_category(fname::Symbol) "symmetry" : "bookkeeping" end +# timer sections are controlled by the `timeit_debug_enabled` switch of the module using `@cached` +_timeit_expr(label, ex) = + Expr(:macrocall, GlobalRef(TimerOutputs, Symbol("@timeit_debug")), LineNumberNode(@__LINE__, @__FILE__), GLOBAL_TIMER, label, ex) + macro cached(ex) - Meta.isexpr(ex, :function) || - error("cached macro can only be used on function definitions") + return _cached(__module__, ex) +end + +function _cached_error(msg, ex) + ex isa Expr && (ex = Base.remove_linenums!(copy(ex))) + Meta.isexpr(ex, :kw) && (ex = Expr(:(=), ex.args...)) # prints as `x = 1` + throw(ArgumentError("@cached: $msg, got `$ex`")) +end + +function _cached(mod::Module, ex) + if !Meta.isexpr(ex, :function) + Meta.isexpr(ex, :(=)) && Meta.isexpr(ex.args[1], (:call, :where, :(::))) && + _cached_error("short-form function definitions are not supported, use `function ... end`", ex.args[1]) + _cached_error("expected a function definition", ex) + end fcall = ex.args[1] if Meta.isexpr(fcall, :where) hasparams = true @@ -72,12 +89,23 @@ macro cached(ex) else typed = false end - Meta.isexpr(fcall, :call) || - error("cached macro can only be used on function definitions") + Meta.isexpr(fcall, :call) || _cached_error("expected a function signature", fcall) fname = fcall.args[1] + # qualified names such as `TensorKit.treebraider` add methods to a function of another module + basename = Meta.isexpr(fname, :.) ? fname.args[end].value : fname + basename isa Symbol || + _cached_error("expected a function name or a qualified name such as `Module.f`", fname) + # the arguments form the cache key, so they must all be named positional arguments + for arg in fcall.args[2:end] + Meta.isexpr(arg, :parameters) && _cached_error("keyword arguments are not supported", fcall) + Meta.isexpr(arg, :kw) && _cached_error("default values are not supported", arg) + Meta.isexpr(arg, :...) && _cached_error("varargs are not supported", arg) + (arg isa Symbol || (Meta.isexpr(arg, :(::)) && length(arg.args) == 2)) || + _cached_error("all arguments must be named, since they form the cache key", arg) + end # timer labels for the cache lookup and the miss-path construction - lookuplabel = string("bookkeeping: cache ", fname) - misslabel = string(_cached_category(fname), ": compute ", fname) + lookuplabel = string("bookkeeping: cache ", basename) + misslabel = string(_cached_category(basename), ": compute ", basename) fargs = fcall.args[2:end] fargnames = map(fargs) do arg if Meta.isexpr(arg, :(::)) @@ -89,7 +117,7 @@ macro cached(ex) _fbody = ex.args[2] # actual implenetation, with underscore name - _fname = Symbol(:_, fname) + _fname = Symbol(:_, basename) _fcall = Expr(:call, _fname, fargs...) if hasparams _fcall = Expr(:where, _fcall, params...) @@ -103,7 +131,7 @@ macro cached(ex) end cachestylevar = gensym(:cachestyle) cachestyleex = Expr( - :(=), cachestylevar, Expr(:call, :CacheStyle, fname, fargnames...) + :(=), cachestylevar, Expr(:call, GlobalRef(@__MODULE__, :CacheStyle), fname, fargnames...) ) newfbody = Expr( :block, cachestyleex, Expr(:call, fname, fargnames..., cachestylevar) @@ -111,11 +139,11 @@ macro cached(ex) newfex = Expr(:function, newfcall, newfbody) # nocache implementation - fnocachecall = Expr(:call, fname, fargs..., :(::NoCache)) + fnocachecall = Expr(:call, fname, fargs..., :(::$NoCache)) if hasparams fnocachecall = Expr(:where, fnocachecall, params...) end - fnocachebody = :(@timeit_debug GLOBAL_TIMER $misslabel $(Expr(:call, _fname, fargnames...))) + fnocachebody = _timeit_expr(misslabel, Expr(:call, _fname, fargnames...)) if typed T = gensym(:T) fnocachebody = Expr(:block, Expr(:(=), T, typeex), Expr(:(::), fnocachebody, T)) @@ -124,16 +152,17 @@ macro cached(ex) # tasklocal cache implementation Dvar = gensym(:D) - flocalcachecall = Expr(:call, fname, fargs..., :(::TaskLocalCache{$Dvar})) + flocalcachecall = Expr(:call, fname, fargs..., :(::$TaskLocalCache{$Dvar})) if hasparams flocalcachecall = Expr(:where, flocalcachecall, params..., Dvar) else flocalcachecall = Expr(:where, flocalcachecall, Dvar) end - localcachename = Symbol(:_tasklocal_, fname, :_cache) + globalcachename = Symbol(:GLOBAL_, uppercase(string(basename)), :_CACHE) + localcachename = Symbol(:_tasklocal_, globalcachename) cachevar = gensym(:cache) getlocalcacheex = :( - $cachevar::$Dvar = get!(task_local_storage(), $localcachename) do + $cachevar::$Dvar = get!(task_local_storage(), $(QuoteNode(localcachename))) do return $Dvar() end ) @@ -143,11 +172,8 @@ macro cached(ex) else key = Expr(:tuple, fargnames...) end - getvalex = :( - @timeit_debug GLOBAL_TIMER $lookuplabel get!($cachevar, $key) do - return @timeit_debug GLOBAL_TIMER $misslabel $_fname($(fargnames...)) - end - ) + missex = _timeit_expr(misslabel, Expr(:call, _fname, fargnames...)) + getvalex = _timeit_expr(lookuplabel, :(get!(() -> $missex, $cachevar, $key))) if typed T = gensym(:T) flocalcachebody = Expr( @@ -168,11 +194,10 @@ macro cached(ex) flocalcacheex = Expr(:function, flocalcachecall, flocalcachebody) # # global cache implementation - fglobalcachecall = Expr(:call, fname, fargs..., :(::GlobalLRUCache)) + fglobalcachecall = Expr(:call, fname, fargs..., :(::$GlobalLRUCache)) if hasparams fglobalcachecall = Expr(:where, fglobalcachecall, params...) end - globalcachename = Symbol(:GLOBAL_, uppercase(string(fname)), :_CACHE) getglobalcachex = Expr(:(=), cachevar, globalcachename) if typed T = gensym(:T) @@ -194,10 +219,12 @@ macro cached(ex) fglobalcacheex = Expr(:function, fglobalcachecall, fglobalcachebody) fglobalcachedef = Expr( :const, - Expr(:(=), globalcachename, :(LRU{Any, Any}(; maxsize = DEFAULT_GLOBALCACHE_SIZE[]))) + Expr(:(=), globalcachename, :($LRU{Any, Any}(; maxsize = $DEFAULT_GLOBALCACHE_SIZE[]))) ) + # caches of other modules (e.g. extensions adding methods) are registered with their module + registername = mod === (@__MODULE__) ? globalcachename : Symbol(nameof(mod), ".", globalcachename) fglobalcacheregister = Expr( - :call, :push!, :GLOBAL_CACHES, :($(QuoteNode(globalcachename)) => $globalcachename) + :call, :push!, GLOBAL_CACHES, :($(QuoteNode(registername)) => $globalcachename) ) # # total expression diff --git a/src/tensors/indexmanipulations.jl b/src/tensors/indexmanipulations.jl index 106f020a4..66f7eb292 100644 --- a/src/tensors/indexmanipulations.jl +++ b/src/tensors/indexmanipulations.jl @@ -657,8 +657,8 @@ function add_transform_kernel!( ) bufsize = buffersize(transformer) if bufsize == 0 # no recoupling needed: every block consists of a single tree - taskforeach(transformer.data, ntasks) do (U, inds_dst, inds_src) - _add_transform_block!(dst, src, p, conjsrc, U, inds_dst, inds_src, nothing, α, β, backend, allocator) + taskforeach(transformer.data, ntasks) do blk + _add_transform_block!(dst, src, p, conjsrc, blk, nothing, α, β, backend, allocator) end else # One max-sized workspace per task (a single one that is reused by all blocks when @@ -669,8 +669,8 @@ function add_transform_kernel!( TO.tensoralloc(storagetype(dst), bufsize, Val(true), allocator) for _ in 1:min(length(transformer.data), ntasks) ] - taskforeach(transformer.data, buffers) do (U, inds_dst, inds_src), buffer - _add_transform_block!(dst, src, p, conjsrc, U, inds_dst, inds_src, buffer, α, β, backend, allocator) + taskforeach(transformer.data, buffers) do blk, buffer + _add_transform_block!(dst, src, p, conjsrc, blk, buffer, α, β, backend, allocator) end foreach(Base.Fix2(TO.tensorfree!, allocator), buffers) TO.allocator_reset!(allocator, cp) @@ -678,48 +678,91 @@ function add_transform_kernel!( return nothing end -# `U` is either a scalar coefficient with integer positions (unique fusion), or a recoupling -# matrix with vectors of positions (generic). function _add_transform_block!( + dst::TransformSubblocks, src::TransformSubblocks, p, conjsrc::Bool, + (coeff, idst, isrc)::UniqueTransformerData, buffer, α, β, backend, allocator + ) + return _add_single_block!(dst, src, p, conjsrc, coeff, idst, isrc, α, β, backend, allocator) +end +function _add_transform_block!( + dst::TransformSubblocks, src::TransformSubblocks, p, conjsrc::Bool, + blk::RecouplingBlock, buffer, α, β, backend, allocator + ) + U = blk.U + return U isa Number ? + _add_single_block!(dst, src, p, conjsrc, U, only(blk.inds_dst), only(blk.inds_src), α, β, backend, allocator) : + _add_recoupling_block!(dst, src, p, conjsrc, U, blk.inds_dst, blk.inds_src, buffer, α, β, backend, allocator) +end + +# single tree: no matmul needed +function _add_single_block!( + dst::TransformSubblocks, src::TransformSubblocks, p, conjsrc::Bool, coeff, idst::Int, isrc::Int, + α, β, backend, allocator + ) + @timeit_debug GLOBAL_TIMER "dense: tensoradd" @inbounds TO.tensoradd!( + dst[idst], src[isrc], p, conjsrc, α * coeff, β, backend, allocator + ) + return nothing +end + +# multi-tree block: pack → recoupling matmul → unpack +function _add_recoupling_block!( dst::TransformSubblocks, src::TransformSubblocks, p, conjsrc::Bool, U, inds_dst, inds_src, buffer, α, β, backend, allocator ) - if length(U) == 1 # single tree: no matmul needed - @timeit_debug GLOBAL_TIMER "dense: tensoradd" @inbounds TO.tensoradd!( - dst[only(inds_dst)], src[only(inds_src)], p, conjsrc, α * only(U), β, backend, allocator + rows, cols = size(U) + sz_src = size(@inbounds(src[first(inds_src)])) + blocksize = prod(sz_src) + ptriv = (ntuple(identity, length(sz_src)), ()) + buffer_dst = StridedView(buffer, (blocksize, rows), (1, blocksize), 0) + buffer_src = StridedView(buffer, (blocksize, cols), (1, blocksize), blocksize * rows) + + # 1. Extract: copy each source block into column i of buffer_src as a flat vector, + # using a trivial permutation so the layout is canonical before the matmul. + @timeit_debug GLOBAL_TIMER "dense: pack" @inbounds for (i, isrc) in enumerate(inds_src) + TO.tensoradd!( + sreshape(view(buffer_src, :, i), sz_src), src[isrc], + ptriv, conjsrc, One(), Zero(), backend, allocator ) - else # Multi-tree block: pack → recoupling matmul → unpack. - rows, cols = size(U) - sz_src = size(@inbounds(src[first(inds_src)])) - blocksize = prod(sz_src) - ptriv = (ntuple(identity, length(sz_src)), ()) - buffer_dst = StridedView(buffer, (blocksize, rows), (1, blocksize), 0) - buffer_src = StridedView(buffer, (blocksize, cols), (1, blocksize), blocksize * rows) - - # 1. Extract: copy each source block into column i of buffer_src as a flat vector, - # using a trivial permutation so the layout is canonical before the matmul. - @timeit_debug GLOBAL_TIMER "dense: pack" @inbounds for (i, isrc) in enumerate(inds_src) - TO.tensoradd!( - sreshape(view(buffer_src, :, i), sz_src), src[isrc], - ptriv, conjsrc, One(), Zero(), backend, allocator - ) - end + end - # 2. Recoupling: buffer_dst = α * buffer_src * U^T (each output tree is a linear - # combination of input trees weighted by the recoupling coefficients). - @timeit_debug GLOBAL_TIMER "dense: recouple mul!" begin - U′ = _adapt_recoupling(storagetype(dst), U) - mul!(buffer_dst, buffer_src, transpose(U′), α, Zero()) - end + # 2. Recoupling: buffer_dst = buffer_src * U^T (each output tree is a linear + # combination of input trees weighted by the recoupling coefficients). + @timeit_debug GLOBAL_TIMER "dense: recouple mul!" _recouple!(buffer, buffer_dst, buffer_src, U) - # 3. Insert: scatter column j of buffer_dst into the destination, applying the - # actual index permutation p in the same tensoradd! call. - @timeit_debug GLOBAL_TIMER "dense: unpack" @inbounds for (j, idst) in enumerate(inds_dst) - TO.tensoradd!( - dst[idst], sreshape(view(buffer_dst, :, j), sz_src), - p, false, One(), β, backend, allocator - ) - end + # 3. Insert: scatter column j of buffer_dst into the destination, applying the + # actual index permutation p and the scaling α in the same tensoradd! call. + @timeit_debug GLOBAL_TIMER "dense: unpack" @inbounds for (j, idst) in enumerate(inds_dst) + TO.tensoradd!( + dst[idst], sreshape(view(buffer_dst, :, j), sz_src), + p, false, α, β, backend, allocator + ) end return nothing end + +# computes `buffer_dst = buffer_src * transpose(U)`, where both are column-major views into `buffer` +_recouple!(buffer, buffer_dst, buffer_src, U) = + (mul!(buffer_dst, buffer_src, transpose(StridedView(U))); nothing) +# real coefficients acting on complex data: recouple the real and imaginary parts in a single real gemm +function _recouple!( + buffer::DenseVector{Complex{R}}, buffer_dst::StridedView, buffer_src::StridedView, U::AbstractMatrix{R} + ) where {R <: LinearAlgebra.BlasReal} + rbuffer = reinterpret(R, buffer) + mul!(_realview(rbuffer, buffer_dst), _realview(rbuffer, buffer_src), transpose(U)) + return nothing +end +# on the CPU, views of reinterpreted arrays are not `StridedMatrix`, so `mul!` would not use BLAS +function _recouple!( + buffer::CPUStorage{Complex{R}}, buffer_dst::StridedView, buffer_src::StridedView, U::Matrix{R} + ) where {R <: LinearAlgebra.BlasReal} + rbuffer = reinterpret(R, buffer) + LinearAlgebra.BLAS.gemm!('N', 'T', one(R), _realview(rbuffer, buffer_src), U, zero(R), _realview(rbuffer, buffer_dst)) + return nothing +end +# real view of the complex column-major view `b` into the reinterpreted buffer `rbuffer` +function _realview(rbuffer::AbstractVector, b::StridedView) + rows, cols = size(b) + offset = 2 * b.offset + return reshape(view(rbuffer, (offset + 1):(offset + 2 * rows * cols)), 2 * rows, cols) +end diff --git a/src/tensors/treetransformers.jl b/src/tensors/treetransformers.jl index 231171da8..2b1b028c1 100644 --- a/src/tensors/treetransformers.jl +++ b/src/tensors/treetransformers.jl @@ -12,6 +12,31 @@ of `tsrc` itself; when `conjsrc` is `true` the fusion trees that are transformed """ abstract type TreeTransformer end +# `StridedSubblocks` report their storage as the `StridedView` parent type, which is `Memory` there +@static if isdefined(Core, :Memory) + const CPUStorage{T} = Union{Array{T}, Memory{T}} +else + const CPUStorage{T} = Array{T} +end + +""" + recoupling_scalartype(A::Type{<:AbstractVector}, Tₛ::Type{<:Number}) -> Type{<:Number} + +Scalar type used to store the recoupling coefficients with sector scalar type `Tₛ` in the +transformers for destination tensors with storagetype `A`. For storage with BLAS scalars, this is +the precision of the storage, where real coefficients are kept real also for complex storage, such +that they can be applied to the real and imaginary parts at once. +""" +function recoupling_scalartype(::Type{A}, ::Type{Tₛ}) where {A, Tₛ <: Number} + T = scalartype(A) + T <: BlasFloat || return Tₛ + return Tₛ <: Real ? real(T) : complex(T) +end + +# matrix type with scalar type `T` in the same storage as the vector type `A` +Base.@assume_effects :foldable recoupling_matrixtype(::Type{A}, ::Type{T}) where {A, T} = + Core.Compiler.return_type(similar, Tuple{A, Type{T}, Dims{2}}) + # (coefficient, destination position, source position) const UniqueTransformerData{T} = Tuple{T, Int, Int} @@ -29,37 +54,51 @@ struct UniqueTreeTransformer{T, N} <: TreeTransformer structure_src::Vector{StridedStructure{N}} end -# (recoupling matrix, destination positions, source positions): U[j, i] maps source i onto destination j -const GenericTransformerData{T} = Tuple{Matrix{T}, Vector{Int}, Vector{Int}} +""" + RecouplingBlock{T, M} +Recoupling data of a single [`FusionTreeBlock`](@ref), mapping the subblocks at the source +positions `inds_src` onto those at the destination positions `inds_dst`. For blocks consisting of a +single tree, `U::T` is a scalar coefficient, while for other blocks it is a recoupling matrix `U::M`, +where `U[j, i]` maps source `i` onto destination `j`. """ - GenericTreeTransformer{T, N} <: TreeTransformer +struct RecouplingBlock{T, M <: AbstractMatrix{T}} + U::Union{T, M} + inds_dst::Vector{Int} + inds_src::Vector{Int} + # `M` cannot be inferred from a scalar `U`, so the parameters are always specified + RecouplingBlock{T, M}(U, inds_dst, inds_src) where {T, M <: AbstractMatrix{T}} = + new{T, M}(U, inds_dst, inds_src) +end + +""" + GenericTreeTransformer{T, M, N} <: TreeTransformer Tree transformation for sectors with multiple fusion channels, where the subblocks of a [`FusionTreeBlock`](@ref) map onto the subblocks of the transformed block through a recoupling -matrix, stored as `(U, inds_dst, inds_src)`. The subblock structures of the destination and -source spaces are kept alongside, such that the [`StridedSubblocks`](@ref) of both tensors can be -created without further lookups. +matrix, stored as a [`RecouplingBlock{T, M}`](@ref RecouplingBlock). The subblock structures of the +destination and source spaces are kept alongside, such that the [`StridedSubblocks`](@ref) of both +tensors can be created without further lookups. """ -struct GenericTreeTransformer{T, N} <: TreeTransformer - data::Vector{GenericTransformerData{T}} +struct GenericTreeTransformer{T, M <: AbstractMatrix{T}, N} <: TreeTransformer + data::Vector{RecouplingBlock{T, M}} structure_dst::Vector{StridedStructure{N}} structure_src::Vector{StridedStructure{N}} end -function UniqueTreeTransformer(transform, p, Vdst, Vsrc, conjsrc::Bool) +function UniqueTreeTransformer(::Type{A}, transform, p, Vdst, Vsrc, conjsrc::Bool) where {A} t₀ = Base.time() spacecheck_transform(permute, Vdst, Vsrc, p, conjsrc) src_trees, dst_trees = fusiontrees(Vsrc), fusiontrees(Vdst) - T = sectorscalartype(sectortype(Vdst)) + T = recoupling_scalartype(A, sectorscalartype(sectortype(Vdst))) data = Vector{UniqueTransformerData{T}}(undef, length(src_trees)) @timeit_debug GLOBAL_TIMER "symmetry: tree transform" for (isrc, (f₁, f₂)) in enumerate(src_trees) f_dst, coeff = transform(conjsrc ? (f₂, f₁) : (f₁, f₂)) _, (_, idst) = gettoken(dst_trees, f_dst) - data[isrc] = (coeff, idst, isrc) + data[isrc] = (convert(T, coeff), idst, isrc) end structure_dst = degeneracystructure(Vdst).subblockstructure @@ -72,7 +111,7 @@ function UniqueTreeTransformer(transform, p, Vdst, Vsrc, conjsrc::Bool) return transformer end -function GenericTreeTransformer(transform, p, Vdst, Vsrc, conjsrc::Bool) +function GenericTreeTransformer(::Type{A}, transform, p, Vdst, Vsrc, conjsrc::Bool) where {A} t₀ = Base.time() spacecheck_transform(permute, Vdst, Vsrc, p, conjsrc) # the fusion blocks that are transformed are those of the adjoint space for a conjugated source @@ -80,18 +119,19 @@ function GenericTreeTransformer(transform, p, Vdst, Vsrc, conjsrc::Bool) src_trees, dst_trees = fusiontrees(Vsrc), fusiontrees(Vdst) structure_dst = degeneracystructure(Vdst).subblockstructure structure_src = degeneracystructure(Vsrc).subblockstructure - T = sectorscalartype(sectortype(Vsrc)) + T = recoupling_scalartype(A, sectorscalartype(sectortype(Vsrc))) + M = recoupling_matrixtype(A, T) fblocks = @timeit_debug GLOBAL_TIMER "bookkeeping: fusionblocks" fusionblocks(Vsrc′) nblocks = length(fblocks) - data = Vector{GenericTransformerData{T}}(undef, nblocks) + data = Vector{RecouplingBlock{T, M}}(undef, nblocks) weights = Vector{Int}(undef, nblocks) nthreads = get_num_manipulation_threads() @timeit_debug GLOBAL_TIMER "symmetry: recoupling matrices" begin taskforeach(1:nblocks, nthreads) do i fs_src = fblocks[i] - fs_dst, U = transform(fs_src) + fs_dst, U₀ = transform(fs_src) @timeit_debug GLOBAL_TIMER "bookkeeping: subblock positions" begin # the token into the fusion tree `Indices` is the subblock position inds_src = map(fusiontrees(fs_src)) do (f₁, f₂) @@ -103,27 +143,28 @@ function GenericTreeTransformer(transform, p, Vdst, Vsrc, conjsrc::Bool) return idst end end - data[i] = (U, inds_dst, inds_src) + U = length(U₀) == 1 ? convert(T, only(U₀)) : convert(M, U₀) + data[i] = RecouplingBlock{T, M}(U, inds_dst, inds_src) # cost model: L input blocks each going to L output blocks of a given length - weights[i] = length(U) * prod(structure_dst[first(inds_dst)][1]) + weights[i] = length(U₀) * prod(structure_dst[first(inds_dst)][1]) @debug( lazy"Created recoupling block for uncoupled: $(fs_src.uncoupled)", - sz = size(U), sparsity = count(!iszero, U) / length(U) + sz = size(U₀), sparsity = count(!iszero, U₀) / length(U₀) ) end end # sort by (approximate) weight to facilitate multi-threading strategies @timeit_debug GLOBAL_TIMER "bookkeeping: sort" Base.permute!(data, sortperm(weights; rev = true)) - transformer = GenericTreeTransformer(data, structure_dst, structure_src) + transformer = GenericTreeTransformer{T, M, numind(Vdst)}(data, structure_dst, structure_src) Δt = Base.time() - t₀ @debug( lazy"TreeTransformer for $Vsrc to $Vdst via $p", conjsrc, nblocks = nblocks, - sz_median = nblocks > 0 ? size(data[cld(end, 2)][1], 1) : 0, - sz_max = nblocks > 0 ? size(data[1][1], 1) : 0, + sz_median = nblocks > 0 ? length(data[cld(end, 2)].inds_dst) : 0, + sz_max = nblocks > 0 ? length(data[1].inds_dst) : 0, Δt ) @@ -140,63 +181,55 @@ the recoupling matrix. buffersize(::UniqueTreeTransformer) = 0 function buffersize(transformer::GenericTreeTransformer) structure_src = transformer.structure_src - return maximum(transformer.data; init = 0) do (U, _, inds_src) - return length(U) == 1 ? 0 : prod(structure_src[first(inds_src)][1]) * sum(size(U)) + return maximum(transformer.data; init = 0) do blk + blk.U isa Number && return 0 + return prod(structure_src[first(blk.inds_src)][1]) * sum(size(blk.U)) end end -function treetransformertype(Vdst, Vsrc) +function treetransformertype(::Type{A}, Vdst, Vsrc) where {A} I = sectortype(Vdst) - T = sectorscalartype(I) + T = recoupling_scalartype(A, sectorscalartype(I)) N = numind(Vdst) - return FusionStyle(I) == UniqueFusion() ? UniqueTreeTransformer{T, N} : GenericTreeTransformer{T, N} + FusionStyle(I) == UniqueFusion() && return UniqueTreeTransformer{T, N} + return GenericTreeTransformer{T, recoupling_matrixtype(A, T), N} end function TreeTransformer( - transform::Function, p, Vdst::HomSpace{S}, Vsrc::HomSpace{S}, conjsrc::Bool - ) where {S} + ::Type{A}, transform::Function, p, Vdst::HomSpace{S}, Vsrc::HomSpace{S}, conjsrc::Bool + ) where {A, S} I = sectortype(Vdst) return FusionStyle(I) == UniqueFusion() ? - UniqueTreeTransformer(transform, p, Vdst, Vsrc, conjsrc) : - GenericTreeTransformer(transform, p, Vdst, Vsrc, conjsrc) + UniqueTreeTransformer(A, transform, p, Vdst, Vsrc, conjsrc) : + GenericTreeTransformer(A, transform, p, Vdst, Vsrc, conjsrc) end # braid is special because it has levels function treebraider( tdst::AbstractTensorMap, tsrc::AbstractTensorMap, p::Index2Tuple, conjsrc::Bool, levels::IndexTuple ) - return treebraider(space(tdst), space(tsrc), p, conjsrc, levels) + return treebraider(storagetype(tdst), space(tdst), space(tsrc), p, conjsrc, levels) end @cached function treebraider( - Vdst::TensorMapSpace, Vsrc::TensorMapSpace, p::Index2Tuple, conjsrc::Bool, levels::IndexTuple - )::treetransformertype(Vdst, Vsrc) + A::Type{TA}, Vdst::TensorMapSpace, Vsrc::TensorMapSpace, p::Index2Tuple, conjsrc::Bool, levels::IndexTuple + )::treetransformertype(A, Vdst, Vsrc) where {TA} Vsrc′, p′ = conjsrc ? (Vsrc', adjointtensorindices(Vsrc, p)) : (Vsrc, p) # levels are attached to the legs, so they follow the same relabeling as the permutation levels′ = conjsrc ? TupleTools.getindices(levels, adjointtensorindices(Vsrc′, allind(Vsrc′))) : levels levels″ = (TupleTools.getindices(levels′, codomainind(Vsrc′)), TupleTools.getindices(levels′, domainind(Vsrc′))) fusiontreebraider(f) = braid(f, p′, levels″) - return TreeTransformer(fusiontreebraider, p, Vdst, Vsrc, conjsrc) + return TreeTransformer(A, fusiontreebraider, p, Vdst, Vsrc, conjsrc) end function treetransposer(tdst::AbstractTensorMap, tsrc::AbstractTensorMap, p::Index2Tuple, conjsrc::Bool) - return treetransposer(space(tdst), space(tsrc), p, conjsrc) + return treetransposer(storagetype(tdst), space(tdst), space(tsrc), p, conjsrc) end @cached function treetransposer( - Vdst::TensorMapSpace, Vsrc::TensorMapSpace, p::Index2Tuple, conjsrc::Bool - )::treetransformertype(Vdst, Vsrc) + A::Type{TA}, Vdst::TensorMapSpace, Vsrc::TensorMapSpace, p::Index2Tuple, conjsrc::Bool + )::treetransformertype(A, Vdst, Vsrc) where {TA} p′ = conjsrc ? adjointtensorindices(Vsrc, p) : p fusiontreetransform(f) = transpose(f, p′) - return TreeTransformer(fusiontreetransform, p, Vdst, Vsrc, conjsrc) + return TreeTransformer(A, fusiontreetransform, p, Vdst, Vsrc, conjsrc) end # default cachestyle is GlobalLRUCache - -# For CPU arrays the recoupling matrix can be used as is, also when the scalar types -# do not match, since Strided handles mixed-eltype mul! without the copy that -# Adapt.adapt would make (which additionally dispatches dynamically). Other storage -# types (e.g. GPU arrays) do require the conversion. -# TODO: transformers with dedicated storagetypes -# `StridedSubblocks` report their storage as the `StridedView` parent type, which is `Memory` there -const CPUStorage = @static isdefined(Core, :Memory) ? Union{Array, Memory} : Array -_adapt_recoupling(::Type{<:CPUStorage}, U::Matrix) = StridedView(U) -_adapt_recoupling(::Type{A}, U::Matrix) where {A} = Adapt.adapt(A, StridedView(U))