Skip to content
Draft
Show file tree
Hide file tree
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
22 changes: 17 additions & 5 deletions ext/ReactantCUDAExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,12 @@ using Reactant.Ops: @opcall
using Enzyme
using Adapt: Adapt, adapt
using CUDA: CUDA, CuDim, DenseCuArray, unsafe_cached_load
# Compatibility for CUDA v5 and v6
@static if isdefined(CUDA, :_derived_array)
# Extend CUDA's internal derived-array helper so traced arrays follow the
# same reshape/reinterpret reconstruction path when it is available; the
# CuTracedArray-specific method is defined alongside reshape below.
import CUDA: _derived_array
end

const CUVERSION = isdefined(CUDA, :CUDACore) ? 6 : 5
if CUVERSION == 6
Expand Down Expand Up @@ -59,6 +64,9 @@ struct CuTracedArray{T,N,A,Size} <: DenseArray{T,N}
ptr = Base.reinterpret(Core.LLVMPtr{T,CUDA.AS.Global}, Base.pointer_from_objref(xs))
return new(ptr)
end
function CuTracedArray{T,N,A,Size}(ptr::Core.LLVMPtr{T,A}) where {T,N,A,Size}
return new(ptr)
end
end

Reactant.use_overlayed_version(::CuTracedArray) = true
Expand Down Expand Up @@ -562,15 +570,13 @@ function Base.reinterpret(::Type{T}, a::CuTracedArray{S,N,A}) where {T,S,N,A}
err === nothing || throw(err)

if sizeof(T) == sizeof(S) # fast case
return CuTracedArray{T,N,A}(
reinterpret(Core.LLVMPtr{T,A}, a.ptr), size(a), a.maxsize
)
return _derived_array(a, T, size(a))
end

isize = size(a)
size1 = div(isize[1] * sizeof(S), sizeof(T))
osize = tuple(size1, Base.tail(isize)...)
return CuTracedArray{T,N,A}(reinterpret(Core.LLVMPtr{T,A}, a.ptr), osize, a.maxsize)
return _derived_array(a, T, osize)
end

## reshape
Expand All @@ -589,6 +595,12 @@ function Base.reshape(a::CuTracedArray{T,M,A}, dims::NTuple{N,Int}) where {T,N,M
return _derived_array(a, T, dims)
end

@inline function _derived_array(
a::CuTracedArray{<:Any,<:Any,A}, ::Type{T}, osize::Dims{N}
) where {T,N,A}
return CuTracedArray{T,N,A,osize}(reinterpret(Core.LLVMPtr{T,A}, a.ptr))
end

struct ReactantKernelAdaptor end

function Adapt.adapt_storage(to::ReactantKernelAdaptor, p::CUDA.CuPtr)
Expand Down
19 changes: 19 additions & 0 deletions test/integration/cuda.jl
Original file line number Diff line number Diff line change
Expand Up @@ -114,6 +114,25 @@ function sin!(x, y)
return nothing
end

function reshape_kernel!(out, x)
xr = reshape(x, 2, 2)
@inbounds out[1] = xr[1, 2]
return nothing
end

function reshape!(out, x)
@cuda blocks = 1 threads = 1 reshape_kernel!(out, x)
return nothing
end

@testset "Reshape Kernel" begin
x = Reactant.to_rarray([1.0, 2.0, 3.0, 4.0])
out = Reactant.to_rarray(zeros(Float64, 1))

@jit reshape!(out, x)
@test Array(out) == [3.0]
end

@testset "Sin Kernel" begin
oA = collect(Float64, 1:1:64)
A = Reactant.to_rarray(oA)
Expand Down