Skip to content
Merged
207 changes: 205 additions & 2 deletions ext/TensorKitGPUArraysExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Comment thread
kshyatt marked this conversation as resolved.

# 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))
Comment thread
kshyatt marked this conversation as resolved.
return nothing
end

end
Loading