Skip to content

storagetype-adapted treetransformers - #556

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

lkdvos wants to merge 8 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
Comment thread src/tensors/indexmanipulations.jl Outdated
Comment thread src/tensors/indexmanipulations.jl Outdated
Comment thread src/tensors/indexmanipulations.jl
Comment thread src/tensors/treetransformers.jl Outdated
Comment thread src/tensors/treetransformers.jl
Comment thread src/tensors/treetransformers.jl
Comment thread src/tensors/treetransformers.jl Outdated
lkdvos and others added 7 commits October 5, 2026 09:11
`@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>

@lkdvos lkdvos left a comment

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

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

I've rebased this on otp of the GPU kernels and attempted to simplify a bit.

Comment thread src/tensors/treetransformers.jl Outdated
Comment thread src/tensors/treetransformers.jl
Comment thread src/tensors/treetransformers.jl
Comment thread src/tensors/indexmanipulations.jl Outdated
Comment thread src/tensors/indexmanipulations.jl Outdated
Comment thread src/tensors/indexmanipulations.jl
@lkdvos
lkdvos requested a review from kshyatt October 5, 2026 16:44
Comment on lines +130 to +134
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₀)

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 a specific reason to make this U₀ ? Also U is defined at this point, no?


# 3. Insert: scatter column j of buffer_dst into the destination, applying the
# actual index permutation p in the same tensoradd! call.
# actual index permutation p and the scaling α in the same tensoradd! call.

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.

Out of curiosity: I think this makes more sense, but was there a practical or performance reason for shifting where α is applied?

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))
length(U) == 1 && return 0
return prod(structure_src[first(inds_src)][1]) * sum(size(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.

Any reason for this change? I don't think this really adds to readability.

Comment on lines +175 to +176
FusionStyle(I) == UniqueFusion() && return UniqueTreeTransformer{T, N}
return GenericTreeTransformer{T, N}

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.

Same question here.

return idst
end
end
U = convert(Matrix{T}, 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.

I would have thought that the storage type information would have been used to also immediately convert this to the proper type of array, i.e. host vs device/gpu?

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.

3 participants