@@ -10,6 +10,7 @@ from libc.string cimport memset, memcpy
1010
1111from cuda.bindings cimport cydriver
1212
13+ from cuda.core.experimental._stream cimport Stream as cyStream
1314from 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+
289304class 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:
0 commit comments