Skip to content
158 changes: 154 additions & 4 deletions ext/TensorKitGPUArraysExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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}}}}
Expand Down Expand Up @@ -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
Expand All @@ -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 unique (one-tree) block of a
`GenericTreeTransformer`.
Comment thread
kshyatt marked this conversation as resolved.
"""
struct UniqueTransformerBlock{T, N}
coeff::T
Expand All @@ -167,6 +167,34 @@ 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`.
"""
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}}}
unique_blocks::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))...)

Expand Down Expand Up @@ -198,11 +226,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}
unique_blocks = UniqueTransformerBlock{T, N}[]
blocks = GenericTransformerBlock{N}[]
coeffs = T[]
structs_dst = TreeStructure{N}[]
structs_src = TreeStructure{N}[]

(; structure_dst, structure_src) = transformer
Comment thread
kshyatt marked this conversation as resolved.
for (U, inds_dst, inds_src) in transformer.data
if length(U) == 1 # same as the unique (Abelian) case
push!(
unique_blocks, _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

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(unique_blocks, unique_offsets, unique_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}
Expand All @@ -220,6 +295,14 @@ function Adapt.adapt_structure(to, t::DeviceUniqueTreeTransformer)
)
end

function Adapt.adapt_structure(to, t::DeviceGenericTreeTransformer)
return DeviceGenericTreeTransformer(
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)
)
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:
Expand Down Expand Up @@ -252,6 +335,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
Expand Down Expand Up @@ -302,6 +386,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)))
Comment thread
lkdvos marked this conversation as resolved.
@inbounds for j in 1:blk.cols
# TODO is there a more efficient way to do this read?
coeff = coeffs[blk.u_offset + 1 + i + (j - 1) * blk.rows]
iszero(coeff) && continue
pst_src, offs_src = structs_src[blk.src_offset + j]
acc += coeff * 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
Expand All @@ -312,6 +437,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!(
Expand All @@ -325,4 +461,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.unique_blocks, α, β, Val(N))
_launch_generic!(dst.data, src.data, op, device, α, β, Val(N))
return nothing
end

end
Loading