diff --git a/ext/TensorKitGPUArraysExt.jl b/ext/TensorKitGPUArraysExt.jl index 158890aa8..3176336b1 100644 --- a/ext/TensorKitGPUArraysExt.jl +++ b/ext/TensorKitGPUArraysExt.jl @@ -3,14 +3,18 @@ module TensorKitGPUArraysExt using GPUArrays using GPUArrays: @allowscalar using GPUArrays.KernelAbstractions: @kernel, @index, get_backend - +using Adapt +using TensorKit.LRUCache: LRU +using TensorKit.TupleTools using Strided: StridedViews using MatrixAlgebraKit, Adapt using TensorKit +using TensorKit.TensorOperations: linearize, DefaultAllocator using TensorKit.Factorizations using TensorKit.Factorizations: AbstractAlgorithm using TensorKit: SectorDict, tensormaptype, scalar, similarstoragetype, AdjointTensorMap, scalartype, project_symmetric_and_check -import TensorKit: randisometry, rand, randn, fill_braidingsubblock! +using TensorKit: StridedSubblocks, UniqueTreeTransformer +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}}}} # COV_EXCL_START @@ -121,5 +125,204 @@ function TensorKit.scalar(t::TensorMap{T, S, 0, 0, <:AnyGPUArray}) where {T, S} return isempty(inds) ? zero(scalartype(t)) : @allowscalar @inbounds t.data[only(inds)] end +# Device-side tree transformers +# ----------------------------- +# A `TreeTransformer` stores the mapping between subblock positions (plus recoupling +# coefficients) and the `StridedStructure`s of the source and destination spaces. +# We resolve the positions into sizes/strides/offsets and pack all the information on the +# CPU side into dense vectors of numbers, plus some accounting information so we know how +# to unpack in the kernel. Also, we can permute the source strides "in advance" on the CPU side. +# We also precompute the strides of the subblock each kernel index will work on, +# 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 +# - Use a 2D grid for the Generic case + +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`. +""" +struct UniqueTransformerBlock{T, N} + coeff::T + sz::NTuple{N, Int} + dense_strides::NTuple{N, Int} + strides_dst::NTuple{N, Int} + offsets_dst::Int + permuted_strides_src::NTuple{N, Int} # source strides, permuted by `p` + offsets_src::Int +end + +# device-side simple struct that GPU kernels can use +struct DeviceUniqueTreeTransformer{VB <: AbstractVector{<:UniqueTransformerBlock}, VO <: AbstractVector{Int}} + blocks::VB + work_offsets::VO + nwork::Int +end + +# strides of a dense array of shape `sz` +_dense_strides(size::Dims) = (1, Base.front(cumprod(size))...) + +# `permute(Vsrc, p) == Vdst` is enforced when the transformer is built, so the permuted +# source shape always matches `sz_dst` and the two views share Cartesian inds. +function _unique_block( + coeff::T, (size_dst, strides_dst, offsets_dst), (_, strides_src, offsets_src), p + ) where {T} + 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) + offsets = cumsum(work) + pushfirst!(offsets, 0) + total = pop!(offsets) + return offsets, total +end + +function DeviceUniqueTreeTransformer(transformer::UniqueTreeTransformer{T, N}, p) where {T, N} + (; structure_dst, structure_src) = transformer + blocks = UniqueTransformerBlock{T, N}[ + _unique_block(coeff, structure_dst[idst], structure_src[isrc], p) + for (coeff, idst, isrc) in transformer.data + ] + work_offsets, nwork = _work_offsets(prod(blk.sz) for blk in blocks) + return DeviceUniqueTreeTransformer(blocks, work_offsets, nwork) +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` +to `Float64`, while `adapt(StoreAdaptor(proto), array)` does not. +""" +struct StorageAdaptor{A <: AbstractArray} + proto::A +end +function Adapt.adapt_storage(a::StorageAdaptor, x::AbstractArray) + dst = similar(a.proto, eltype(x), size(x)) + isempty(x) && return dst + return copy!(dst, x) +end + +function Adapt.adapt_structure(to, t::DeviceUniqueTreeTransformer) + return DeviceUniqueTreeTransformer( + Adapt.adapt(to, t.blocks), Adapt.adapt(to, t.work_offsets), t.nwork + ) +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: +# - transformer which is an immutable struct. It's hashed and compared by the identity of +# its fields, which avoids walking every recoupling matrix on every lookup. +# Holding it in the key also keeps it alive, so its `objectid` cannot be +# reused by a different transformer while the entry is cached. +# - the storage type +# - `p`, which is baked into the permuted source strides. +# TODO: should this live in the main package? +const DEVICE_TRANSFORMER_CACHE = LRU{Any, Any}(; maxsize = TensorKit.DEFAULT_GLOBALCACHE_SIZE[]) + +function __init__() + push!(TensorKit.GLOBAL_CACHES, :DEVICE_TRANSFORMER_CACHE => DEVICE_TRANSFORMER_CACHE) + return nothing +end + +# We have this complicated setup because a naive `adapt` doesn't work. +# Rather we copy everything to GPU-native arrays and have kernels that can work +# with that. +function device_transformer(proto::AbstractArray, transformer, p) + key = (transformer, typeof(proto), p) + return get!(DEVICE_TRANSFORMER_CACHE, key) do + # be careful about the lifetime of these, since they live in a global cache and + # thus persist beyond the call + GPUArrays.@uncached Adapt.adapt( + StorageAdaptor(proto), _device_transformer(transformer, p) + ) + end +end + +_device_transformer(t::UniqueTreeTransformer, p) = DeviceUniqueTreeTransformer(t, p) + +# COV_EXCL_START +# kernels are not reachable by coverage + +# largest `i` with `offsets[i] <= w`. This corresponds to the +# block which this kernel thread will work on. Since this is +# used inside a GPU kernel, searchsortedlast/searchsortedfirst +# won't work. +@inline function _searchblock(offsets, w) + lo, hi = 1, length(offsets) + while lo < hi + mid = (lo + hi + 1) >>> 1 + if @inbounds offsets[mid] <= w + lo = mid + else + hi = mid - 1 + end + end + return lo +end + +# Cartesian coordinates of the `w`-th (0-based) entry of a dense subblock of shape `sz`. +@inline function _coordinates(w, size::NTuple{N, Int}, dense_strides::NTuple{N, Int}) where {N} + return ntuple(n -> (w ÷ dense_strides[n]) % size[n], Val(N)) +end + +# finds the overall linear index in the output and input arrays corresponding to the **sublock** +# coordinates currently being worked on +@inline function _linear_index(coords::NTuple{N, Int}, strides::NTuple{N, Int}, offset) where {N} + return offset + sum(ntuple(n -> coords[n] * strides[n], Val(N))) + 1 +end + +# One thread per destination element in `data_dst`. `op` is `identity` or `conj`, and is +# applied to the source data only (not to the coefficients). +@kernel function unique_batched_permute_kernel!( + data_dst, data_src, op, blocks, work_offsets, α, β, nwork, ::Val{N} + ) where {N} + w = @index(Global, Linear) - 1 + if w < nwork + b = _searchblock(work_offsets, w) + blk = @inbounds blocks[b] + coords = _coordinates(w - (@inbounds work_offsets[b]), blk.sz, blk.dense_strides) + i_dst = _linear_index(coords, blk.strides_dst, blk.offsets_dst) + i_src = _linear_index(coords, blk.permuted_strides_src, blk.offsets_src) + @inbounds data_dst[i_dst] = α * blk.coeff * op(data_src[i_src]) + β * data_dst[i_dst] + end +end + +# COV_EXCL_STOP + +function _launch_unique!(data_dst, data_src, op, transformer, α, β, ::Val{N}) where {N} + nwork = transformer.nwork + nwork == 0 && return nothing + unique_batched_permute_kernel!(get_backend(data_dst))( + data_dst, data_src, op, transformer.blocks, transformer.work_offsets, α, β, nwork, + Val(N); ndrange = nwork + ) + return nothing +end + +const GPUStridedSubblocks = StridedSubblocks{<:AnyGPUArray} + +function TensorKit.add_transform_kernel!( + dst::GPUStridedSubblocks, src::GPUStridedSubblocks, p, conjsrc::Bool, + transformer::UniqueTreeTransformer{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 + _launch_unique!(dst.data, src.data, op, device, α, β, Val(N)) + return nothing +end end