Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
159 changes: 155 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 degenerate (one-tree) block of a
`GenericTreeTransformer`.
"""
struct UniqueTransformerBlock{T, N}
coeff::T
Expand All @@ -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))...)

Expand Down Expand Up @@ -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}
Expand All @@ -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:
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)))
Comment thread
lkdvos marked this conversation as resolved.
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
Expand All @@ -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!(
Expand All @@ -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
Loading