Skip to content

storagetype-adapted treetransformers - #556

Open
lkdvos wants to merge 3 commits into
mainfrom
ld-storagetype
Open

lkdvos wants to merge 3 commits into
mainfrom
ld-storagetype

Conversation

@lkdvos

@lkdvos lkdvos commented Sep 24, 2026 •

Copy link
Copy Markdown
Member

This is preparatory work to hopefully simplify some of the logic in #533:

The goal is to make it possible and convenient to have a dispatched method for caching a different kind of treetransformer for device vs host code. This required 3 changes:

  1. The @cached macro needed a bit of hygiene to easily work from modules that aren't TensorKit (such as extensions)
  2. The treebraider calls etc now take the storagetype of the tdst as a first argument to facilitate dispatch
  3. The GenericTreeTransformer is already slightly improved: it now stores the unitary recoupling coefficients in the correct storagetype to avoid data transfer.
  4. Bonus: I used the complex * real matrix trick that reinterprets it as a real * real matrix with twice the rows, thereby leveraging BLAS even for real recoupling coefficients and complex tensors (eg. SU(2))!

This should allow the GPU extension to simply overload these functions and return custom structs.

@lkdvos
lkdvos requested a review from kshyatt September 24, 2026 18:59
@codecov

codecov Bot commented Sep 25, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 92.00000% with 8 lines in your changes missing coverage. Please review.

Files with missing lines Patch % Lines
src/auxiliary/caches.jl 81.25% 6 Missing ⚠️
src/tensors/indexmanipulations.jl 97.36% 1 Missing ⚠️
src/tensors/treetransformers.jl 96.66% 1 Missing ⚠️
Files with missing lines Coverage Δ
src/tensors/indexmanipulations.jl 91.76% <97.36%> (+0.62%) ⬆️
src/tensors/treetransformers.jl 95.55% <96.66%> (-0.79%) ⬇️
src/auxiliary/caches.jl 86.11% <81.25%> (-2.90%) ⬇️

... and 2 files with indirect coverage changes

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@lkdvos
lkdvos marked this pull request as ready for review September 25, 2026 17:57
@lkdvos
lkdvos requested a review from Jutho September 25, 2026 17:57
@lkdvos

lkdvos commented Sep 25, 2026

Copy link
Copy Markdown
Member Author

Benchmarks for the complex, non-abelian case with real recoupling coefficients: the SU2Irrep entries of the indexmanipulations (permute) and tensornetworks (MPO/PEPO/MERA contraction) suites, run with ComplexF64. Single-threaded (Julia and BLAS), on one rusty rome node, comparing 3b285456 (base) with this branch before the rebase onto #554, which only touches AD rules.

group time ratio new/base notes
PEPO (8 cases) 0.73 – 0.90 largest gains for the largest bond dimensions, e.g. [8, 2, 4, 200]: 30.0 s → 21.8 s
MERA (6 cases) 0.88 – 0.97 improves with size, D = 28: 9.81 s → 8.66 s
MPO (8 cases) 0.94 – 1.00 dominated by the dense contractions themselves
permute [512, 512] 0.99 no recoupling involved
permute [48, 48, 48, 48] (4 cases) 1.01 – 1.04 see below

The small difference for the 4-leg permutes does not reproduce locally, where the branch is 2–6% faster for the same cases. These are dominated by the single-tree blocks (29 of 37 blocks), which go through the unchanged permuting tensoradd! path and are memory-bound, while the recoupling step itself, around 10% of the runtime here, becomes about 2× faster with the real gemm. I therefore attribute this difference to process-to-process variation on the cluster nodes.

Full results
benchmark                                                                                                    base          new     time   memory  judgement
indexmanipulations / permute / ("ComplexF64", "SU2Irrep", "[48, 48, 48, 48]", "[1.0, 1.0, 1.0, 1.0]", "Any[[1, 3], [2, 4]]")   888.069 μs   917.715 μs    1.033    0.992  invariant
indexmanipulations / permute / ("ComplexF64", "SU2Irrep", "[48, 48, 48, 48]", "[1.0, 1.0, 1.0, 1.0]", "Any[[1, 3], [2, 4]]", "adjoint")     1.039 ms     1.046 ms    1.006    0.992  invariant
indexmanipulations / permute / ("ComplexF64", "SU2Irrep", "[48, 48, 48, 48]", "[1.0, 1.0, 1.0, 1.0]", "Any[[4, 2, 3], [1]]")   903.287 μs   934.827 μs    1.035    0.992  invariant
indexmanipulations / permute / ("ComplexF64", "SU2Irrep", "[48, 48, 48, 48]", "[1.0, 1.0, 1.0, 1.0]", "Any[[4, 2, 3], [1]]", "adjoint")   842.422 μs   870.415 μs    1.033    0.992  invariant
indexmanipulations / permute / ("ComplexF64", "SU2Irrep", "[512, 512]", "[1.0, 1.0]", "Any[[2, 1], Any[]]")    50.796 μs    50.155 μs    0.987    1.000  invariant
tensornetworks / mera / ("ComplexF64", "SU2Irrep", 4, 2.0)                                              19.150 ms    17.890 ms    0.934    0.928  invariant
tensornetworks / mera / ("ComplexF64", "SU2Irrep", 8, 2.0)                                             244.712 ms   236.644 ms    0.967    0.958  invariant
tensornetworks / mera / ("ComplexF64", "SU2Irrep", 12, 2.0)                                            295.401 ms   278.342 ms    0.942    0.966  invariant
tensornetworks / mera / ("ComplexF64", "SU2Irrep", 16, 2.0)                                               2.425 s      2.320 s    0.957    0.979  invariant
tensornetworks / mera / ("ComplexF64", "SU2Irrep", 22, 2.0)                                               4.375 s      3.982 s    0.910    0.990  invariant
tensornetworks / mera / ("ComplexF64", "SU2Irrep", 28, 2.0)                                               9.810 s      8.656 s    0.882    0.996  invariant
tensornetworks / mpo / ("ComplexF64", "SU2Irrep", "[160, 5, 3]", 2)                                    464.077 μs   445.572 μs    0.960    0.990  invariant
tensornetworks / mpo / ("ComplexF64", "SU2Irrep", "[200, 20, 20]", 2)                                    9.220 ms     8.778 ms    0.952    0.988  invariant
tensornetworks / mpo / ("ComplexF64", "SU2Irrep", "[2560, 5, 3]", 2)                                   286.971 ms   286.479 ms    0.998    1.000  invariant
tensornetworks / mpo / ("ComplexF64", "SU2Irrep", "[40, 5, 3]", 2)                                     221.097 μs   208.655 μs    0.944    0.950  invariant
tensornetworks / mpo / ("ComplexF64", "SU2Irrep", "[400, 20, 20]", 2)                                   37.247 ms    36.762 ms    0.987    0.997  invariant
tensornetworks / mpo / ("ComplexF64", "SU2Irrep", "[400, 40, 40]", 2)                                  166.384 ms   165.472 ms    0.995    0.998  invariant
tensornetworks / mpo / ("ComplexF64", "SU2Irrep", "[6120, 5, 3]", 2)                                      3.566 s      3.562 s    0.999    1.000  invariant
tensornetworks / mpo / ("ComplexF64", "SU2Irrep", "[640, 5, 3]", 2)                                      6.682 ms     6.634 ms    0.993    0.999  invariant
tensornetworks / pepo / ("ComplexF64", "SU2Irrep", "[10, 2, 2, 50]", 2.0)                                 5.308 s      4.790 s    0.902    0.987  invariant
tensornetworks / pepo / ("ComplexF64", "SU2Irrep", "[10, 3, 2, 100]", 2.0)                                6.712 s      5.725 s    0.853    0.991  invariant
tensornetworks / pepo / ("ComplexF64", "SU2Irrep", "[4, 2, 2, 100]", 2.0)                              984.190 ms   873.143 ms    0.887    0.984  invariant
tensornetworks / pepo / ("ComplexF64", "SU2Irrep", "[4, 4, 4, 200]", 2.0)                                14.122 s     10.456 s    0.740    0.995  improvement
tensornetworks / pepo / ("ComplexF64", "SU2Irrep", "[6, 2, 2, 100]", 2.0)                                 1.174 s      1.034 s    0.880    0.989  invariant
tensornetworks / pepo / ("ComplexF64", "SU2Irrep", "[6, 3, 4, 200]", 2.0)                                 4.361 s      3.500 s    0.803    0.996  improvement
tensornetworks / pepo / ("ComplexF64", "SU2Irrep", "[8, 2, 2, 100]", 2.0)                                 6.975 s      5.688 s    0.816    0.991  improvement
tensornetworks / pepo / ("ComplexF64", "SU2Irrep", "[8, 2, 4, 200]", 2.0)                                30.014 s     21.827 s    0.727    0.997  improvement

Comment thread src/auxiliary/caches.jl Outdated
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 || error("cached macro can only be used on function definitions")

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.

any way we could embed the basename in here to give people more info?

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

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

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

Comment thread src/tensors/treetransformers.jl Outdated
"""
struct RecouplingBlock{T, M <: AbstractMatrix{T}}
coeff::T
U::Union{Nothing, M}

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.

Why not make the type of U a type parameter?


# `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}}

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 name seems weird to me. Maybe it should be DefaultStorage or something? It's defined in opposition to the GPU extensions which seems strange

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

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.

What about storage that isn't made of BLAS scalars?

Comment thread src/tensors/treetransformers.jl Outdated
`@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 <noreply@anthropic.com>
`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 <noreply@anthropic.com>
…orage

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 <noreply@anthropic.com>

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants