From e575048e39c1dbd5324d100a1e3a7652253e4698 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 17 Aug 2026 13:11:37 -0400 Subject: [PATCH 1/9] Implement different kernel approach for GPU-side braiding --- ext/TensorKitGPUArraysExt.jl | 154 +++++++++++++++++++++++++++++++++-- 1 file changed, 149 insertions(+), 5 deletions(-) diff --git a/ext/TensorKitGPUArraysExt.jl b/ext/TensorKitGPUArraysExt.jl index 3176336b1..8ba3d40d1 100644 --- a/ext/TensorKitGPUArraysExt.jl +++ b/ext/TensorKitGPUArraysExt.jl @@ -13,7 +13,7 @@ using TensorKit.TensorOperations: linearize, DefaultAllocator using TensorKit.Factorizations using TensorKit.Factorizations: AbstractAlgorithm using TensorKit: SectorDict, tensormaptype, scalar, similarstoragetype, AdjointTensorMap, scalartype, project_symmetric_and_check -using TensorKit: StridedSubblocks, UniqueTreeTransformer +using TensorKit: StridedSubblocks, UniqueTreeTransformer, GenericTreeTransformer import TensorKit: randisometry, rand, randn, fill_braidingsubblock!, add_transform_kernel! function TensorKit.fill_braidingsubblock!(data::TD, val) where {T, TD <: Union{<:AnyGPUMatrix{T}, <:StridedViews.StridedView{T, 4, <:AnyGPUArray{T}}}} @@ -136,7 +136,6 @@ end # so that the GPU thread can recover the Cartesian coordinates it will need for input/ouput, # and a running `work_offsets` count of destination elements, so that a kernel can run one # thread per output element and recover which input that element belongs with. - # Some possible TODO here: # - Try cuTILE as this is a classic tile programming problem # - Use shared memory to coalesce the reads @@ -148,7 +147,8 @@ const TreeStructure{N} = Tuple{NTuple{N, Int}, Int} UniqueTransformerBlock{T, N} `isbits` descriptor for a subblock that is a *single* scaled permutation: -an entry of a `UniqueTreeTransformer`. +an entry of a `UniqueTreeTransformer` or a degenerate (one-tree) block of a +`GenericTreeTransformer`. """ struct UniqueTransformerBlock{T, N} coeff::T @@ -167,6 +167,35 @@ struct DeviceUniqueTreeTransformer{VB <: AbstractVector{<:UniqueTransformerBlock nwork::Int end +""" + GenericTransformerBlock{N} + +Descriptor for a recoupling block of a `GenericTreeTransformer`, indexing into the +flat `coeffs`/`structs_dst`/`structs_src` vectors of a `DeviceGenericTreeTransformer`. +** All offsets are 0-based since it makes the arithmetic easier. ** +""" +struct GenericTransformerBlock{N} + sz::NTuple{N, Int} + densestrides::NTuple{N, Int} + rows::Int + cols::Int + u_offset::Int # location in the flattened U vector to find this block's U + dst_offset::Int + src_offset::Int +end + +# force all the type signatures here to make sure doing something wrong fails +# before the kernel launch. Kernel error dumps are awful and hard to interpret. +struct DeviceGenericTreeTransformer{VO <: AbstractVector{Int}, DA <: DeviceAbelianTreeTransformer{<:Any, VO}, VB <: AbstractVector{<:GenericTransformerBlock}, VC <: AbstractVector{<:Number}, VS <: AbstractVector{<:Tuple{<:Tuple{Vararg{Int}}, Int}}} + degenerate::DA # length(U) = 1 blocks, can be handled by Abelian kernel + blocks::VB + work_offsets::VO + nwork::Int + coeffs::VC # every `U`, concatenated in column-major order + structs_dst::VS + structs_src::VS +end + # strides of a dense array of shape `sz` _dense_strides(size::Dims) = (1, Base.front(cumprod(size))...) @@ -178,7 +207,7 @@ function _unique_block( return UniqueTransformerBlock{T, length(size_dst)}( coeff, size_dst, _dense_strides(size_dst), strides_dst, offsets_dst, TupleTools.getindices(strides_src, p), offsets_src - ) + ) end function _work_offsets(work) @@ -198,11 +227,51 @@ function DeviceUniqueTreeTransformer(transformer::UniqueTreeTransformer{T, N}, p return DeviceUniqueTreeTransformer(blocks, work_offsets, nwork) end +function DeviceGenericTreeTransformer( + transformer::GenericTreeTransformer{T, N}, p + ) where {T, N} + degenerate = UniqueTransformerBlock{T, N}[] + blocks = GenericTransformerBlock{N}[] + coeffs = T[] + structs_dst = TreeStructure{N}[] + structs_src = TreeStructure{N}[] + + for (U, (size_dst, strides_dst), (size_src, strides_src)) in transformer.data + if length(U) == 1 # same as the unique (Abelian) case + push!( + degenerate, _abelian_block( + only(U), (size_dst, only(strides_dst)...), (size_src, only(strides_src)...), p + ) + ) + else + push!( + blocks, GenericTransformerBlock{N}( + size_dst, _dense_strides(size_dst), size(U, 1), size(U, 2), + length(coeffs), length(structs_dst), length(structs_src) + ) + ) + append!(coeffs, U) + append!(structs_dst, strides_dst) + for (stride_src, offset_src) in strides_src + push!(structs_src, (_permutestrides(stride_src, p), offset_src)) + end + end + end + + degenerate_offsets, degenerate_nwork = _work_offsets(prod(blk.sz) for blk in degenerate) + work_offsets, nwork = _work_offsets(blk.rows * prod(blk.sz) for blk in blocks) + return DeviceGenericTreeTransformer( + DeviceUniqueTreeTransformer(degenerate, degenerate_offsets, degenerate_nwork), + blocks, work_offsets, nwork, coeffs, structs_dst, structs_src + ) +end + """ StorageAdaptor(proto) `Adapt` adaptor moving arrays onto the same device and array type as `proto`, preserving -their element type. For `proto::CuVector{Float64}` and `array::Vector{Int}`, the call `adapt(typeof(proto), array)` would force-convert the element type `Int` +their element type. For `proto::CuVector{Float64}` and `array::Vector{Int}`, +the call `adapt(typeof(proto), array)` would force-convert the element type `Int` to `Float64`, while `adapt(StoreAdaptor(proto), array)` does not. """ struct StorageAdaptor{A <: AbstractArray} @@ -220,6 +289,14 @@ function Adapt.adapt_structure(to, t::DeviceUniqueTreeTransformer) ) end +function Adapt.adapt_structure(to, t::DeviceGenericTreeTransformer) + return DeviceGenericTreeTransformer( + Adapt.adapt(to, t.degenerate), Adapt.adapt(to, t.blocks), + Adapt.adapt(to, t.work_offsets), t.nwork, Adapt.adapt(to, t.coeffs), + Adapt.adapt(to, t.structs_dst), Adapt.adapt(to, t.structs_src) + ) +end + # Copying a transformer to GPU is more expensive than running it, so we cache the device # copy in a global LRU cache, registered in `TensorKit.GLOBAL_CACHES` so that # `empty_globalcaches!` also frees the device memory. The key is: @@ -252,6 +329,7 @@ function device_transformer(proto::AbstractArray, transformer, p) end _device_transformer(t::UniqueTreeTransformer, p) = DeviceUniqueTreeTransformer(t, p) +_device_transformer(t::GenericTreeTransformer, p) = DeviceGenericTreeTransformer(t, p) # COV_EXCL_START # kernels are not reachable by coverage @@ -302,6 +380,47 @@ end # COV_EXCL_STOP +# One thread per destination element in `data_dst`. This makes much better use of the +# GPU "massive parallelism" as compared to the one-thread-per-subtransformer approach. +# It also more evenly divides the work among threads so the work profile is less +# jagged. Unlike the CPU implementation, there is no extract → recouple → insert process: +# BLAS is not generally reachable from inside a kernel, and fusing the recoupling into +# the strided gather lets us remove the buffer entirely. +# TODO: what about symmetries like SU(3), where the column by column approach is not +# optimal? +@kernel function generic_batched_permute_kernel!( + data_dst, data_src, op, blocks, work_offsets, coeffs, structs_dst, structs_src, + α, β, nwork, ::Val{N} + ) where {N} + w = @index(Global, Linear) - 1 + if w < nwork + # bookkeeping to figure out where to read from and write to + b = _searchblock(work_offsets, w) + blk = @inbounds blocks[b] + local_w = w - (@inbounds work_offsets[b]) + blocksize = prod(blk.sz) + i = local_w ÷ blocksize # 0-based destination tree + coords = _coordinates(local_w % blocksize, blk.sz, blk.densestrides) + + st_dst, offs_dst = @inbounds structs_dst[blk.dst_offset + i + 1] + i_dst = _offset(coords, st_dst, offs_dst) + + # dst_i = β * dst_i + α * Σ_j U[i, j] * permute(src_j, p): each output tree is a + # linear combination of the input trees weighted by the recoupling coefficients. + # The permutation of src_j was already done by permuting its strides before the + # kernel launched. + acc = zero(promote_type(eltype(data_src), eltype(coeffs))) + for j in 1:blk.cols + # TODO is there a more efficient way to do this read? + coeff = @inbounds coeffs[blk.u_offset + (j - 1) * blk.rows + i + 1] + iszero(coeff) && continue + pst_src, offs_src = @inbounds structs_src[blk.src_offset + j] + acc += coeff * @inbounds op(data_src[_offset(coords, pst_src, offs_src)]) + end + @inbounds data_dst[i_dst] = α * acc + β * data_dst[i_dst] + end +end + function _launch_unique!(data_dst, data_src, op, transformer, α, β, ::Val{N}) where {N} nwork = transformer.nwork nwork == 0 && return nothing @@ -312,6 +431,17 @@ function _launch_unique!(data_dst, data_src, op, transformer, α, β, ::Val{N}) return nothing end +function _launch_generic!(data_dst, data_src, op, transformer, α, β, ::Val{N}) where {N} + nwork = transformer.nwork + nwork == 0 && return nothing + generic_batched_permute_kernel!(get_backend(data_dst))( + data_dst, data_src, op, transformer.blocks, transformer.work_offsets, + transformer.coeffs, transformer.structs_dst, transformer.structs_src, α, β, nwork, + Val(N); ndrange = nwork + ) + return nothing +end + const GPUStridedSubblocks = StridedSubblocks{<:AnyGPUArray} function TensorKit.add_transform_kernel!( @@ -325,4 +455,18 @@ function TensorKit.add_transform_kernel!( return nothing end +function TensorKit.add_transform_kernel!( + data_dst::GPUStridedSubblocks, data_src::GPUStridedSubblocks, p, conjsrc::Bool, + transformer::GenericTreeTransformer{T, N}, α, β, backend, allocator, ntasks::Int + ) where {T, N} + # GPU-side object to hold the treetransformer information + device = device_transformer(dst.data, transformer, linearize(p)) + op = conjsrc ? conj : identity + # one-tree blocks are a scaled permutation, which the Abelian kernel already handles; the + # two kernels touch disjoint subblocks so the launch order does not matter + _launch_unique!(dst.data, src.data, op, device.degenerate, α, β, Val(N)) + _launch_generic!(dst.data, src.data, op, device, α, β, Val(N)) + return nothing +end + end From b8f607d7e974018e973860f6eb65e3ecf79d3b7e Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 20 Aug 2026 17:06:47 +0200 Subject: [PATCH 2/9] Send tensorfree over to TO --- Project.toml | 3 +++ 1 file changed, 3 insertions(+) diff --git a/Project.toml b/Project.toml index 1f466a811..ea07fb876 100644 --- a/Project.toml +++ b/Project.toml @@ -72,3 +72,6 @@ TimerOutputs = "1" TupleTools = "1.5" VectorInterface = "0.6, 0.7" julia = "1.10" + +[sources] +TensorOperations = {url = "https://github.com/QuantumKitHub/TensorOperations.jl", rev = "ksh/gputensorfree"} From fae44db5ce409e6fa5118bf76e59979ee77f6c9c Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 14 Sep 2026 13:18:23 +0200 Subject: [PATCH 3/9] Fix Project.toml sources --- Project.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/Project.toml b/Project.toml index ea07fb876..4b2db51a9 100644 --- a/Project.toml +++ b/Project.toml @@ -74,4 +74,4 @@ VectorInterface = "0.6, 0.7" julia = "1.10" [sources] -TensorOperations = {url = "https://github.com/QuantumKitHub/TensorOperations.jl", rev = "ksh/gputensorfree"} +TensorOperations = {url = "https://github.com/QuantumKitHub/TensorOperations.jl", rev = "main"} From 5b4f635e07a495504df614cfd1dd6f9a20c2db0f Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 14 Sep 2026 15:17:07 +0200 Subject: [PATCH 4/9] Try to fix ambiguity warning --- ext/TensorKitGPUArraysExt.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ext/TensorKitGPUArraysExt.jl b/ext/TensorKitGPUArraysExt.jl index 8ba3d40d1..52d7d3b93 100644 --- a/ext/TensorKitGPUArraysExt.jl +++ b/ext/TensorKitGPUArraysExt.jl @@ -456,7 +456,7 @@ function TensorKit.add_transform_kernel!( end function TensorKit.add_transform_kernel!( - data_dst::GPUStridedSubblocks, data_src::GPUStridedSubblocks, p, conjsrc::Bool, + dst::GPUStridedSubblocks, src::GPUStridedSubblocks, p, conjsrc::Bool, transformer::GenericTreeTransformer{T, N}, α, β, backend, allocator, ntasks::Int ) where {T, N} # GPU-side object to hold the treetransformer information From d0551094d40d7f99dc8729a82240e13a72f040c2 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 29 Sep 2026 14:32:32 +0200 Subject: [PATCH 5/9] Formatter --- ext/TensorKitGPUArraysExt.jl | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/ext/TensorKitGPUArraysExt.jl b/ext/TensorKitGPUArraysExt.jl index 52d7d3b93..260d0495e 100644 --- a/ext/TensorKitGPUArraysExt.jl +++ b/ext/TensorKitGPUArraysExt.jl @@ -207,7 +207,7 @@ function _unique_block( return UniqueTransformerBlock{T, length(size_dst)}( coeff, size_dst, _dense_strides(size_dst), strides_dst, offsets_dst, TupleTools.getindices(strides_src, p), offsets_src - ) + ) end function _work_offsets(work) @@ -457,7 +457,7 @@ end function TensorKit.add_transform_kernel!( dst::GPUStridedSubblocks, src::GPUStridedSubblocks, p, conjsrc::Bool, - transformer::GenericTreeTransformer{T, N}, α, β, backend, allocator, ntasks::Int + transformer::GenericTreeTransformer{T, N}, α, β, backend, allocator, ntasks::Int ) where {T, N} # GPU-side object to hold the treetransformer information device = device_transformer(dst.data, transformer, linearize(p)) From 17281a55b03b19e43d8b4b4bd2772c7bc386bc09 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 29 Sep 2026 16:16:20 +0200 Subject: [PATCH 6/9] Leftover Abelian from rebase --- ext/TensorKitGPUArraysExt.jl | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/ext/TensorKitGPUArraysExt.jl b/ext/TensorKitGPUArraysExt.jl index 260d0495e..fbc061606 100644 --- a/ext/TensorKitGPUArraysExt.jl +++ b/ext/TensorKitGPUArraysExt.jl @@ -186,8 +186,8 @@ end # force all the type signatures here to make sure doing something wrong fails # before the kernel launch. Kernel error dumps are awful and hard to interpret. -struct DeviceGenericTreeTransformer{VO <: AbstractVector{Int}, DA <: DeviceAbelianTreeTransformer{<:Any, VO}, VB <: AbstractVector{<:GenericTransformerBlock}, VC <: AbstractVector{<:Number}, VS <: AbstractVector{<:Tuple{<:Tuple{Vararg{Int}}, Int}}} - degenerate::DA # length(U) = 1 blocks, can be handled by Abelian kernel +struct DeviceGenericTreeTransformer{VO <: AbstractVector{Int}, DA <: DeviceUniqueTreeTransformer{<:Any, VO}, VB <: AbstractVector{<:GenericTransformerBlock}, VC <: AbstractVector{<:Number}, VS <: AbstractVector{<:Tuple{<:Tuple{Vararg{Int}}, Int}}} + degenerate::DA # length(U) = 1 blocks, can be handled by unique kernel blocks::VB work_offsets::VO nwork::Int @@ -462,7 +462,7 @@ function TensorKit.add_transform_kernel!( # GPU-side object to hold the treetransformer information device = device_transformer(dst.data, transformer, linearize(p)) op = conjsrc ? conj : identity - # one-tree blocks are a scaled permutation, which the Abelian kernel already handles; the + # one-tree blocks are a scaled permutation, which the unique kernel already handles; the # two kernels touch disjoint subblocks so the launch order does not matter _launch_unique!(dst.data, src.data, op, device.degenerate, α, β, Val(N)) _launch_generic!(dst.data, src.data, op, device, α, β, Val(N)) From cb60a2039dee6da41fde90b14de7b9fa470d5b8c Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 29 Sep 2026 13:32:11 -0400 Subject: [PATCH 7/9] Some more rebase cleanups --- Project.toml | 3 --- ext/TensorKitGPUArraysExt.jl | 8 ++++---- 2 files changed, 4 insertions(+), 7 deletions(-) diff --git a/Project.toml b/Project.toml index 4b2db51a9..1f466a811 100644 --- a/Project.toml +++ b/Project.toml @@ -72,6 +72,3 @@ TimerOutputs = "1" TupleTools = "1.5" VectorInterface = "0.6, 0.7" julia = "1.10" - -[sources] -TensorOperations = {url = "https://github.com/QuantumKitHub/TensorOperations.jl", rev = "main"} diff --git a/ext/TensorKitGPUArraysExt.jl b/ext/TensorKitGPUArraysExt.jl index fbc061606..c292f3b4f 100644 --- a/ext/TensorKitGPUArraysExt.jl +++ b/ext/TensorKitGPUArraysExt.jl @@ -239,7 +239,7 @@ function DeviceGenericTreeTransformer( for (U, (size_dst, strides_dst), (size_src, strides_src)) in transformer.data if length(U) == 1 # same as the unique (Abelian) case push!( - degenerate, _abelian_block( + degenerate, _unique_block( only(U), (size_dst, only(strides_dst)...), (size_src, only(strides_src)...), p ) ) @@ -253,7 +253,7 @@ function DeviceGenericTreeTransformer( append!(coeffs, U) append!(structs_dst, strides_dst) for (stride_src, offset_src) in strides_src - push!(structs_src, (_permutestrides(stride_src, p), offset_src)) + push!(structs_src, (TupleTools.getindices(stride_src, p), offset_src)) end end end @@ -403,7 +403,7 @@ end coords = _coordinates(local_w % blocksize, blk.sz, blk.densestrides) st_dst, offs_dst = @inbounds structs_dst[blk.dst_offset + i + 1] - i_dst = _offset(coords, st_dst, offs_dst) + i_dst = _linear_index(coords, st_dst, offs_dst) # dst_i = β * dst_i + α * Σ_j U[i, j] * permute(src_j, p): each output tree is a # linear combination of the input trees weighted by the recoupling coefficients. @@ -415,7 +415,7 @@ end coeff = @inbounds coeffs[blk.u_offset + (j - 1) * blk.rows + i + 1] iszero(coeff) && continue pst_src, offs_src = @inbounds structs_src[blk.src_offset + j] - acc += coeff * @inbounds op(data_src[_offset(coords, pst_src, offs_src)]) + acc += coeff * @inbounds op(data_src[_linear_index(coords, pst_src, offs_src)]) end @inbounds data_dst[i_dst] = α * acc + β * data_dst[i_dst] end From 76dbb8a6df576eb4a41723994c4afa176bd3882b Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 29 Sep 2026 13:44:48 -0400 Subject: [PATCH 8/9] Adapt to some more changes in 526 --- ext/TensorKitGPUArraysExt.jl | 17 ++++++++++++----- 1 file changed, 12 insertions(+), 5 deletions(-) diff --git a/ext/TensorKitGPUArraysExt.jl b/ext/TensorKitGPUArraysExt.jl index c292f3b4f..20f5c5eeb 100644 --- a/ext/TensorKitGPUArraysExt.jl +++ b/ext/TensorKitGPUArraysExt.jl @@ -236,14 +236,17 @@ function DeviceGenericTreeTransformer( structs_dst = TreeStructure{N}[] structs_src = TreeStructure{N}[] - for (U, (size_dst, strides_dst), (size_src, strides_src)) in transformer.data + (; structure_dst, structure_src) = transformer + for (U, inds_dst, inds_src) in transformer.data if length(U) == 1 # same as the unique (Abelian) case push!( degenerate, _unique_block( - only(U), (size_dst, only(strides_dst)...), (size_src, only(strides_src)...), p + only(U), structure_dst[only(inds_dst)], structure_src[only(inds_src)], p ) ) else + # all trees in a block share the same subblock size + size_dst = first(structure_dst[first(inds_dst)]) push!( blocks, GenericTransformerBlock{N}( size_dst, _dense_strides(size_dst), size(U, 1), size(U, 2), @@ -251,9 +254,13 @@ function DeviceGenericTreeTransformer( ) ) append!(coeffs, U) - append!(structs_dst, strides_dst) - for (stride_src, offset_src) in strides_src - push!(structs_src, (TupleTools.getindices(stride_src, p), offset_src)) + for idst in inds_dst + _, strides_dst, offset_dst = structure_dst[idst] + push!(structs_dst, (strides_dst, offset_dst)) + end + for isrc in inds_src + _, strides_src, offset_src = structure_src[isrc] + push!(structs_src, (TupleTools.getindices(strides_src, p), offset_src)) end end end From 803830c5ed564530f5957263762baa1c8e2569e5 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Fri, 2 Oct 2026 03:52:08 -0400 Subject: [PATCH 9/9] Degenerate -> unique and remove extraneous comment --- ext/TensorKitGPUArraysExt.jl | 17 ++++++++--------- 1 file changed, 8 insertions(+), 9 deletions(-) diff --git a/ext/TensorKitGPUArraysExt.jl b/ext/TensorKitGPUArraysExt.jl index 20f5c5eeb..f77c40be4 100644 --- a/ext/TensorKitGPUArraysExt.jl +++ b/ext/TensorKitGPUArraysExt.jl @@ -147,7 +147,7 @@ const TreeStructure{N} = Tuple{NTuple{N, Int}, Int} UniqueTransformerBlock{T, N} `isbits` descriptor for a subblock that is a *single* scaled permutation: -an entry of a `UniqueTreeTransformer` or a degenerate (one-tree) block of a +an entry of a `UniqueTreeTransformer` or a unique (one-tree) block of a `GenericTreeTransformer`. """ struct UniqueTransformerBlock{T, N} @@ -172,7 +172,6 @@ end Descriptor for a recoupling block of a `GenericTreeTransformer`, indexing into the flat `coeffs`/`structs_dst`/`structs_src` vectors of a `DeviceGenericTreeTransformer`. -** All offsets are 0-based since it makes the arithmetic easier. ** """ struct GenericTransformerBlock{N} sz::NTuple{N, Int} @@ -187,7 +186,7 @@ end # force all the type signatures here to make sure doing something wrong fails # before the kernel launch. Kernel error dumps are awful and hard to interpret. struct DeviceGenericTreeTransformer{VO <: AbstractVector{Int}, DA <: DeviceUniqueTreeTransformer{<:Any, VO}, VB <: AbstractVector{<:GenericTransformerBlock}, VC <: AbstractVector{<:Number}, VS <: AbstractVector{<:Tuple{<:Tuple{Vararg{Int}}, Int}}} - degenerate::DA # length(U) = 1 blocks, can be handled by unique kernel + unique_blocks::DA # length(U) = 1 blocks, can be handled by unique kernel blocks::VB work_offsets::VO nwork::Int @@ -230,7 +229,7 @@ end function DeviceGenericTreeTransformer( transformer::GenericTreeTransformer{T, N}, p ) where {T, N} - degenerate = UniqueTransformerBlock{T, N}[] + unique_blocks = UniqueTransformerBlock{T, N}[] blocks = GenericTransformerBlock{N}[] coeffs = T[] structs_dst = TreeStructure{N}[] @@ -240,7 +239,7 @@ function DeviceGenericTreeTransformer( for (U, inds_dst, inds_src) in transformer.data if length(U) == 1 # same as the unique (Abelian) case push!( - degenerate, _unique_block( + unique_blocks, _unique_block( only(U), structure_dst[only(inds_dst)], structure_src[only(inds_src)], p ) ) @@ -265,10 +264,10 @@ function DeviceGenericTreeTransformer( end end - degenerate_offsets, degenerate_nwork = _work_offsets(prod(blk.sz) for blk in degenerate) + unique_offsets, unique_nwork = _work_offsets(prod(blk.sz) for blk in unique_blocks) work_offsets, nwork = _work_offsets(blk.rows * prod(blk.sz) for blk in blocks) return DeviceGenericTreeTransformer( - DeviceUniqueTreeTransformer(degenerate, degenerate_offsets, degenerate_nwork), + DeviceUniqueTreeTransformer(unique_blocks, unique_offsets, unique_nwork), blocks, work_offsets, nwork, coeffs, structs_dst, structs_src ) end @@ -298,7 +297,7 @@ end function Adapt.adapt_structure(to, t::DeviceGenericTreeTransformer) return DeviceGenericTreeTransformer( - Adapt.adapt(to, t.degenerate), Adapt.adapt(to, t.blocks), + Adapt.adapt(to, t.unique_blocks), Adapt.adapt(to, t.blocks), Adapt.adapt(to, t.work_offsets), t.nwork, Adapt.adapt(to, t.coeffs), Adapt.adapt(to, t.structs_dst), Adapt.adapt(to, t.structs_src) ) @@ -471,7 +470,7 @@ function TensorKit.add_transform_kernel!( op = conjsrc ? conj : identity # one-tree blocks are a scaled permutation, which the unique kernel already handles; the # two kernels touch disjoint subblocks so the launch order does not matter - _launch_unique!(dst.data, src.data, op, device.degenerate, α, β, Val(N)) + _launch_unique!(dst.data, src.data, op, device.unique_blocks, α, β, Val(N)) _launch_generic!(dst.data, src.data, op, device, α, β, Val(N)) return nothing end