diff --git a/ext/ReactantCUDAExt.jl b/ext/ReactantCUDAExt.jl index 0b2cd80243..49162c60e4 100644 --- a/ext/ReactantCUDAExt.jl +++ b/ext/ReactantCUDAExt.jl @@ -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 @@ -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 @@ -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 @@ -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) diff --git a/test/integration/cuda.jl b/test/integration/cuda.jl index 9ab811a74f..8336bc13f1 100644 --- a/test/integration/cuda.jl +++ b/test/integration/cuda.jl @@ -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)