Skip to content

Commit 0c64e8e

Browse files
committed
ensure we have C access for DeviceMemoryResource
1 parent 638bf59 commit 0c64e8e

3 files changed

Lines changed: 72 additions & 48 deletions

File tree

cuda_core/cuda/core/experimental/_memory.pyx

Lines changed: 56 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@ from libc.string cimport memset, memcpy
1010

1111
from cuda.bindings cimport cydriver
1212

13+
from cuda.core.experimental._stream cimport Stream as cyStream
1314
from cuda.core.experimental._utils.cuda_utils cimport (
1415
_check_driver_error as raise_if_driver_error,
1516
check_or_create_options,
@@ -95,7 +96,10 @@ cdef class Buffer:
9596
the behavior depends on the underlying memory resource.
9697
"""
9798
if self._ptr and self._mr is not None:
98-
self._mr.deallocate(self._ptr, self._size, stream)
99+
if isinstance(self._mr, _cyMemoryResource):
100+
(<_cyMemoryResource>(self._mr))._deallocate(self._ptr, self._size, <cyStream>stream)
101+
else:
102+
self._mr.deallocate(self._ptr, self._size, stream)
99103
self._ptr = 0
100104
self._mr = None
101105
self._ptr_obj = None
@@ -286,6 +290,17 @@ cdef class Buffer:
286290
return Buffer._init(ptr, size, mr=mr)
287291

288292

293+
cdef class _cyMemoryResource:
294+
"""
295+
Internal only. Responsible for offering fast C method access.
296+
"""
297+
cdef Buffer _allocate(self, size_t size, cyStream stream):
298+
raise NotImplementedError
299+
300+
cdef int _deallocate(self, intptr_t ptr, size_t size, cyStream stream) except?-1:
301+
raise NotImplementedError
302+
303+
289304
class MemoryResource(abc.ABC):
290305
"""Abstract base class for memory resources that manage allocation and deallocation of buffers.
291306
@@ -542,33 +557,7 @@ class DeviceMemoryResourceAttributes:
542557
_ipc_registry = {}
543558

544559

545-
cdef class _DeviceMemoryResourceBase:
546-
"""Internal only. Responsible for offering C layout & attribute access."""
547-
cdef:
548-
int _dev_id
549-
cydriver.CUmemoryPool _mempool_handle
550-
object _attributes
551-
cydriver.CUmemAllocationHandleType _ipc_handle_type
552-
bint _mempool_owned
553-
bint _is_mapped
554-
object _uuid
555-
IPCAllocationHandle _alloc_handle
556-
557-
def __cinit__(self):
558-
self._dev_id = cydriver.CU_DEVICE_INVALID
559-
self._mempool_handle = NULL
560-
self._attributes = None
561-
self._ipc_handle_type = cydriver.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_MAX
562-
self._mempool_owned = False
563-
self._is_mapped = False
564-
self._uuid = None
565-
self._alloc_handle = None
566-
567-
def __dealloc__(self):
568-
pass
569-
570-
571-
cdef class DeviceMemoryResource(_DeviceMemoryResourceBase, MemoryResource):
560+
cdef class DeviceMemoryResource(_cyMemoryResource, MemoryResource):
572561
"""
573562
Create a device memory resource managing a stream-ordered memory pool.
574563
@@ -647,9 +636,27 @@ cdef class DeviceMemoryResource(_DeviceMemoryResourceBase, MemoryResource):
647636
associated MMR.
648637
"""
649638
cdef:
650-
dict __dict__
639+
int _dev_id
640+
cydriver.CUmemoryPool _mempool_handle
641+
object _attributes
642+
cydriver.CUmemAllocationHandleType _ipc_handle_type
643+
bint _mempool_owned
644+
bint _is_mapped
645+
object _uuid
646+
IPCAllocationHandle _alloc_handle
647+
dict __dict__ # TODO: check if we still need this
651648
object __weakref__
652649

650+
def __cinit__(self):
651+
self._dev_id = cydriver.CU_DEVICE_INVALID
652+
self._mempool_handle = NULL
653+
self._attributes = None
654+
self._ipc_handle_type = cydriver.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_MAX
655+
self._mempool_owned = False
656+
self._is_mapped = False
657+
self._uuid = None
658+
self._alloc_handle = None
659+
653660
def __init__(self, device_id: int | Device, options=None):
654661
cdef int dev_id = getattr(device_id, 'device_id', device_id)
655662
opts = check_or_create_options(
@@ -862,6 +869,18 @@ cdef class DeviceMemoryResource(_DeviceMemoryResourceBase, MemoryResource):
862869
raise
863870
return self._alloc_handle
864871

872+
cdef Buffer _allocate(self, size_t size, cyStream stream):
873+
cdef cydriver.CUstream s = stream._handle
874+
cdef cydriver.CUdeviceptr devptr
875+
with nogil:
876+
HANDLE_RETURN(cydriver.cuMemAllocFromPoolAsync(&devptr, size, self._mempool_handle, s))
877+
cdef Buffer buf = Buffer.__new__(Buffer)
878+
buf._ptr = <intptr_t>(devptr)
879+
buf._ptr_obj = None
880+
buf._size = size
881+
buf._mr = self
882+
return buf
883+
865884
def allocate(self, size_t size, stream: Stream = None) -> Buffer:
866885
"""Allocate a buffer of the requested size.
867886

@@ -883,11 +902,14 @@ cdef class DeviceMemoryResource(_DeviceMemoryResourceBase, MemoryResource):
883902
raise TypeError("Cannot allocate from a mapped IPC-enabled memory resource")
884903
if stream is None:
885904
stream = default_stream()
886-
cdef cydriver.CUstream s = <cydriver.CUstream><uintptr_t>(stream.handle)
887-
cdef cydriver.CUdeviceptr devptr
905+
return self._allocate(size, <cyStream>stream)
906+
907+
cdef int _deallocate(self, intptr_t ptr, size_t size, cyStream stream) except?-1:
908+
cdef cydriver.CUstream s = stream._handle
909+
cdef cydriver.CUdeviceptr devptr = <cydriver.CUdeviceptr>ptr
888910
with nogil:
889-
HANDLE_RETURN(cydriver.cuMemAllocFromPoolAsync(&devptr, size, self._mempool_handle, s))
890-
return Buffer._init(<intptr_t>devptr, size, self)
911+
HANDLE_RETURN(cydriver.cuMemFreeAsync(devptr, s))
912+
return 0
891913

892914
def deallocate(self, ptr: DevicePointerT, size_t size, stream: Stream = None):
893915
"""Deallocate a buffer previously allocated by this resource.
@@ -904,10 +926,7 @@ cdef class DeviceMemoryResource(_DeviceMemoryResourceBase, MemoryResource):
904926
"""
905927
if stream is None:
906928
stream = default_stream()
907-
cdef cydriver.CUstream s = <cydriver.CUstream><uintptr_t>(stream.handle)
908-
cdef cydriver.CUdeviceptr devptr = <cydriver.CUdeviceptr><intptr_t>ptr
909-
with nogil:
910-
HANDLE_RETURN(cydriver.cuMemFreeAsync(devptr, s))
929+
self._deallocate(<intptr_t>ptr, size, <cyStream>stream)
911930

912931
@property
913932
def attributes(self) -> DeviceMemoryResourceAttributes:

cuda_core/cuda/core/experimental/_stream.pxd

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,3 +6,19 @@ from cuda.bindings cimport cydriver
66

77

88
cdef cydriver.CUstream _try_to_get_stream_ptr(obj: IsStreamT) except*
9+
10+
11+
cdef class Stream:
12+
13+
cdef:
14+
cydriver.CUstream _handle
15+
object _owner
16+
object _builtin
17+
object _nonblocking
18+
object _priority
19+
cydriver.CUdevice _device_id
20+
cydriver.CUcontext _ctx_handle
21+
22+
cpdef close(self)
23+
cdef int _get_context(self) except?-1 nogil
24+
cdef int _get_device_and_context(self) except?-1

cuda_core/cuda/core/experimental/_stream.pyx

Lines changed: 0 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -108,18 +108,7 @@ cdef class Stream:
108108
New streams should instead be created through a :obj:`~_device.Device`
109109
object, or created directly through using an existing handle
110110
using Stream.from_handle().
111-
112111
"""
113-
114-
cdef:
115-
cydriver.CUstream _handle
116-
object _owner
117-
object _builtin
118-
object _nonblocking
119-
object _priority
120-
cydriver.CUdevice _device_id
121-
cydriver.CUcontext _ctx_handle
122-
123112
def __cinit__(self):
124113
self._handle = <cydriver.CUstream>(NULL)
125114
self._device_id = cydriver.CU_DEVICE_INVALID # delayed

0 commit comments

Comments
 (0)