diff --git a/ext/TensorKitGPUArraysExt.jl b/ext/TensorKitGPUArraysExt.jl index 3176336b1..20f5c5eeb 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 <: 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 + 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))...) @@ -198,11 +227,58 @@ 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}[] + + (; 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), 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), + length(coeffs), length(structs_dst), length(structs_src) + ) + ) + append!(coeffs, U) + 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 + + 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 +296,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 +336,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 +387,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 = _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. + # 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[_linear_index(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 +438,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 +462,18 @@ function TensorKit.add_transform_kernel!( return nothing end +function TensorKit.add_transform_kernel!( + 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 + device = device_transformer(dst.data, transformer, linearize(p)) + 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_generic!(dst.data, src.data, op, device, α, β, Val(N)) + return nothing +end + end