diff --git a/src/array.jl b/src/array.jl index c0a183b6..9df190c6 100644 --- a/src/array.jl +++ b/src/array.jl @@ -599,12 +599,22 @@ Base.unsafe_convert(::Type{ZePtr{T}}, A::PermutedDimsArray) where {T} = ## unsafe_wrap """ - unsafe_wrap(Array, arr::oneArray{_,_,oneL0.SharedBuffer}) + unsafe_wrap(Array, arr::oneArray{_,_,<:Union{oneL0.SharedBuffer,oneL0.HostBuffer}}) -Wrap a Julia `Array` around the buffer that backs a `oneArray`. This is only possible if the -GPU array is backed by a shared buffer, i.e. if it was created with `oneArray{T}(undef, ...)`. +Wrap a Julia `Array` around the buffer that backs a `oneArray`, without copying. This is +only possible if the GPU array is backed by memory that is accessible from the host, i.e., +a shared buffer (as created by `oneArray{T}(undef, ...)`) or a host buffer. + +!!! warning + + The returned `Array` does **not** keep `arr` alive. The caller has to keep a reference + to `arr` for as long as the `Array`, or anything derived from it, is used; otherwise + the `Array` may end up referring to freed memory. Device operations execute + asynchronously, so call `synchronize()` before accessing the returned array after + using `arr` on the device. """ -function Base.unsafe_wrap(::Type{Array}, arr::oneArray{T,N,oneL0.SharedBuffer}) where {T,N} +function Base.unsafe_wrap(::Type{Array}, + arr::oneArray{T,N,<:Union{oneL0.SharedBuffer,oneL0.HostBuffer}}) where {T,N} # TODO: can we make this more convenient by increasing the buffer's refcount and using # a finalizer on the Array? does that work when taking views etc of the Array? ptr = reinterpret(Ptr{T}, pointer(arr)) diff --git a/test/array.jl b/test/array.jl index 5d92d759..f71d6ced 100644 --- a/test/array.jl +++ b/test/array.jl @@ -101,6 +101,13 @@ end @test Array(a) == [100, 42] oneAPI.@sync copyto!(a, 2, [200], 1, 1) @test b == [100, 200] + + # the same works for arrays backed by host memory + c = oneVector{Int,oneL0.HostBuffer}([1, 2]) + d = unsafe_wrap(Array, c) + @test d == [1, 2] + d[1] = 100 + @test Array(c) == [100, 2] end # https://github.com/JuliaGPU/CUDA.jl/issues/2191