Skip to content
Merged
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
18 changes: 14 additions & 4 deletions src/array.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down
7 changes: 7 additions & 0 deletions test/array.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading