From 7cc755c9d3a1ada8586e562066b2b4325f4c06e8 Mon Sep 17 00:00:00 2001 From: lkdvos Date: Mon, 28 Sep 2026 13:02:52 -0400 Subject: [PATCH 1/3] Make `@cached` usable from other modules `@cached` escaped its entire expansion, so that the cache styles, the global cache registry, the LRU constructor and the timer were resolved in the calling module, which only works within TensorKit. These are now referenced through `GlobalRef`s, and qualified function names such as `TensorKit.treebraider` are supported, such that package extensions can add cached methods to TensorKit functions. The global cache of such methods lives in the module that defines them, and is registered with its module name so that it is shown separately in `global_cache_info`. Also fixes the task-local cache key, which was spliced in as an identifier instead of as a symbol. Co-Authored-By: Claude Opus 5.5 --- docs/src/Changelog.md | 2 ++ src/auxiliary/caches.jl | 71 ++++++++++++++++++++++++++++------------- 2 files changed, 51 insertions(+), 22 deletions(-) diff --git a/docs/src/Changelog.md b/docs/src/Changelog.md index e3baac971..fe9a09ad4 100644 --- a/docs/src/Changelog.md +++ b/docs/src/Changelog.md @@ -24,6 +24,8 @@ 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) + ### Deprecated ### Removed 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 From 897b8e239b8983febc51838b3cbba96e7b997e0e Mon Sep 17 00:00:00 2001 From: lkdvos Date: Mon, 28 Sep 2026 13:02:53 -0400 Subject: [PATCH 2/3] Construct and cache tree transformers per destination storagetype MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `treebraider` and `treetransposer` now take the storagetype `A` of the destination tensor as their first argument. It is part of the cache key, and it is a dispatch point: a storage-specific method, with its own cache through `@cached`, can return a dedicated transformer type. The recoupling data is stored in the form the kernel needs for `A`: - `recoupling_scalartype` stores the coefficients in the precision of the storage. On CPU, real coefficients stay real for complex data. - Blocks of `GenericTreeTransformer` are stored as `RecouplingBlock`s, a concrete type holding either a host scalar (single tree) or a recoupling matrix in the storage of the destination. For GPU storage the matrices are thus converted once at construction, instead of adapted on every call. In the kernel, `α` is applied in the unpack step, such that the recoupling is a plain matrix product. For complex CPU data with real coefficients, the real and imaginary parts are recoupled in a single real `gemm` on a reinterpreted view of the buffer. Co-Authored-By: Claude Opus 5.5 --- docs/src/Changelog.md | 3 + src/tensors/indexmanipulations.jl | 114 ++++++++++++++++--------- src/tensors/treetransformers.jl | 135 +++++++++++++++++++----------- 3 files changed, 165 insertions(+), 87 deletions(-) diff --git a/docs/src/Changelog.md b/docs/src/Changelog.md index fe9a09ad4..6c662ca1a 100644 --- a/docs/src/Changelog.md +++ b/docs/src/Changelog.md @@ -25,6 +25,7 @@ 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 @@ -34,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/tensors/indexmanipulations.jl b/src/tensors/indexmanipulations.jl index 106f020a4..e988f3dcd 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,84 @@ 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, +# called directly since views of reinterpreted arrays are not `StridedMatrix` and `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..f8000e9a8 100644 --- a/src/tensors/treetransformers.jl +++ b/src/tensors/treetransformers.jl @@ -12,6 +12,37 @@ 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 CPU 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. For other storage with BLAS +scalars, this is the scalar type of the storage. +""" +function recoupling_scalartype(::Type{A}, ::Type{Tₛ}) where {A <: CPUStorage, Tₛ <: Number} + T = scalartype(A) + T <: BlasFloat || return Tₛ + return Tₛ <: Real ? real(T) : complex(T) +end +function recoupling_scalartype(::Type{A}, ::Type{Tₛ}) where {A, Tₛ <: Number} + T = scalartype(A) + T <: BlasFloat || return Tₛ + return Tₛ <: 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 +60,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`. +""" +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, N} <: TreeTransformer + 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 +117,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 +125,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 +149,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 +187,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)) From 56f134c2b060bb59ee9f74069540808c80333871 Mon Sep 17 00:00:00 2001 From: lkdvos Date: Mon, 28 Sep 2026 13:02:53 -0400 Subject: [PATCH 3/3] Recouple complex data with real coefficients in a real gemm on all storage Reinterpreting GPU storage as real yields native GPU arrays, such that the real and imaginary parts of complex data can be recoupled with real coefficients in a single `mul!`, which dispatches to the vendor BLAS. This adds a generic `_recouple!` method for complex `DenseVector` buffers with a real recoupling matrix, keeping the direct `BLAS.gemm!` call only for CPU storage, where views of reinterpreted arrays are not `StridedMatrix`. The CPU rule of `recoupling_scalartype`, keeping real coefficients real, is now used for all storage. Co-Authored-By: Claude Opus 5.5 --- src/tensors/indexmanipulations.jl | 11 +++++++++-- src/tensors/treetransformers.jl | 14 ++++---------- 2 files changed, 13 insertions(+), 12 deletions(-) diff --git a/src/tensors/indexmanipulations.jl b/src/tensors/indexmanipulations.jl index e988f3dcd..66f7eb292 100644 --- a/src/tensors/indexmanipulations.jl +++ b/src/tensors/indexmanipulations.jl @@ -744,8 +744,15 @@ 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, -# called directly since views of reinterpreted arrays are not `StridedMatrix` and `mul!` would not use BLAS +# 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} diff --git a/src/tensors/treetransformers.jl b/src/tensors/treetransformers.jl index f8000e9a8..2b1b028c1 100644 --- a/src/tensors/treetransformers.jl +++ b/src/tensors/treetransformers.jl @@ -23,20 +23,14 @@ 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 CPU 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. For other storage with BLAS -scalars, this is the scalar type of the storage. +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 <: CPUStorage, Tₛ <: Number} - T = scalartype(A) - T <: BlasFloat || return Tₛ - return Tₛ <: Real ? real(T) : complex(T) -end function recoupling_scalartype(::Type{A}, ::Type{Tₛ}) where {A, Tₛ <: Number} T = scalartype(A) T <: BlasFloat || return Tₛ - return Tₛ <: Real ? T : complex(T) + return Tₛ <: Real ? real(T) : complex(T) end # matrix type with scalar type `T` in the same storage as the vector type `A`