Skip to content
Open
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
5 changes: 5 additions & 0 deletions docs/src/Changelog.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down
71 changes: 49 additions & 22 deletions src/auxiliary/caches.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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, :(::))
Expand All @@ -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...)
Expand All @@ -103,19 +131,19 @@ 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)
)
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))
Expand All @@ -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
)
Expand All @@ -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(
Expand All @@ -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)
Expand All @@ -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
Expand Down
121 changes: 82 additions & 39 deletions src/tensors/indexmanipulations.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -669,57 +669,100 @@ 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)
end
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

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

is there no way we could dispatch on this? hugely ugly ternary here

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,

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

this might be a good spot for some ASCII art

# 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

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

this comment seems kind of out of place to me. I don't know if it's needed, or maybe it should be rephrased? It's very Claudey

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
Loading
Loading