From e87a9ee5078ac98a8d762647c9be5eb215674b81 Mon Sep 17 00:00:00 2001 From: "Ralf W. Grosse-Kunstleve" Date: Fri, 14 Aug 2026 11:54:59 -0700 Subject: [PATCH] Avoid aggregate CUmemLocation initialization --- .../core/_memory/_device_memory_resource.pyx | 7 +++--- cuda_core/cuda/core/_memory/_location.pxd | 22 +++++++++---------- .../cuda/core/_memory/_managed_memory_ops.pyx | 5 ++--- cuda_core/cuda/core/_memory/_memory_pool.pyx | 5 +++-- .../cuda/core/_memory/_peer_access_utils.pyx | 7 +++--- cuda_core/cuda/core/graph/_graph_node.pyx | 12 +++++----- 6 files changed, 26 insertions(+), 32 deletions(-) diff --git a/cuda_core/cuda/core/_memory/_device_memory_resource.pyx b/cuda_core/cuda/core/_memory/_device_memory_resource.pyx index 1ee492edd77..7ff17933194 100644 --- a/cuda_core/cuda/core/_memory/_device_memory_resource.pyx +++ b/cuda_core/cuda/core/_memory/_device_memory_resource.pyx @@ -321,10 +321,9 @@ cpdef str DMR_mempool_get_access(DeviceMemoryResource dmr, int device_id): cdef int dev_id = Device(device_id).device_id cdef cydriver.CUmemAccess_flags flags - cdef cydriver.CUmemLocation location = cydriver.CUmemLocation( - type=cydriver.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE, - id=dev_id, - ) + cdef cydriver.CUmemLocation location + location.type = cydriver.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE + location.id = dev_id with nogil: HANDLE_RETURN(cydriver.cuMemPoolGetAccess(&flags, as_cu(dmr._h_pool), &location)) diff --git a/cuda_core/cuda/core/_memory/_location.pxd b/cuda_core/cuda/core/_memory/_location.pxd index e46850ca886..7cee3c6564e 100644 --- a/cuda_core/cuda/core/_memory/_location.pxd +++ b/cuda_core/cuda/core/_memory/_location.pxd @@ -15,24 +15,22 @@ from cuda.bindings cimport cydriver IF CUDA_CORE_BUILD_MAJOR >= 13: cdef inline cydriver.CUmemLocation to_cumemlocation(str kind, int loc_id): + cdef cydriver.CUmemLocation cu_loc if kind == "device": - return cydriver.CUmemLocation( - type=cydriver.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE, - id=loc_id) + cu_loc.type = cydriver.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE + cu_loc.id = loc_id elif kind == "host": - return cydriver.CUmemLocation( - type=cydriver.CUmemLocationType.CU_MEM_LOCATION_TYPE_HOST, - id=0) + cu_loc.type = cydriver.CUmemLocationType.CU_MEM_LOCATION_TYPE_HOST + cu_loc.id = 0 elif kind == "host_numa": - return cydriver.CUmemLocation( - type=cydriver.CUmemLocationType.CU_MEM_LOCATION_TYPE_HOST_NUMA, - id=loc_id) + cu_loc.type = cydriver.CUmemLocationType.CU_MEM_LOCATION_TYPE_HOST_NUMA + cu_loc.id = loc_id elif kind == "host_numa_current": - return cydriver.CUmemLocation( - type=cydriver.CUmemLocationType.CU_MEM_LOCATION_TYPE_HOST_NUMA_CURRENT, - id=0) + cu_loc.type = cydriver.CUmemLocationType.CU_MEM_LOCATION_TYPE_HOST_NUMA_CURRENT + cu_loc.id = 0 else: raise ValueError(f"unknown location kind: {kind!r}") + return cu_loc ELSE: cdef inline cydriver.CUmemLocation to_cumemlocation(str kind, int loc_id): raise NotImplementedError( diff --git a/cuda_core/cuda/core/_memory/_managed_memory_ops.pyx b/cuda_core/cuda/core/_memory/_managed_memory_ops.pyx index dcda07aab06..6d504670f36 100644 --- a/cuda_core/cuda/core/_memory/_managed_memory_ops.pyx +++ b/cuda_core/cuda/core/_memory/_managed_memory_ops.pyx @@ -193,9 +193,8 @@ cdef void _do_single_advise(Buffer buf, object advice_value, object loc, bint al # Driver ignores location for read_mostly / unset_preferred_location # advice values but still validates the CUmemLocation; pass a # host placeholder. - cu_loc = cydriver.CUmemLocation( - type=cydriver.CUmemLocationType.CU_MEM_LOCATION_TYPE_HOST, - id=0) + cu_loc.type = cydriver.CUmemLocationType.CU_MEM_LOCATION_TYPE_HOST + cu_loc.id = 0 else: cu_loc = to_cumemlocation(loc.kind, loc.id) with nogil: diff --git a/cuda_core/cuda/core/_memory/_memory_pool.pyx b/cuda_core/cuda/core/_memory/_memory_pool.pyx index cccc95a01a2..988c3ab532b 100644 --- a/cuda_core/cuda/core/_memory/_memory_pool.pyx +++ b/cuda_core/cuda/core/_memory/_memory_pool.pyx @@ -279,8 +279,9 @@ cdef int MP_init_current_pool( """ IF CUDA_CORE_BUILD_MAJOR >= 13: cdef cydriver.CUmemoryPool pool - cdef cydriver.CUmemLocation loc = cydriver.CUmemLocation( - type=loc_type, id=loc_id) + cdef cydriver.CUmemLocation loc + loc.type = loc_type + loc.id = loc_id with nogil: HANDLE_RETURN(cydriver.cuMemGetMemPool(&pool, &loc, alloc_type)) self._h_pool = create_mempool_handle_ref(pool) diff --git a/cuda_core/cuda/core/_memory/_peer_access_utils.pyx b/cuda_core/cuda/core/_memory/_peer_access_utils.pyx index 69d59f9e005..b39a1838f79 100644 --- a/cuda_core/cuda/core/_memory/_peer_access_utils.pyx +++ b/cuda_core/cuda/core/_memory/_peer_access_utils.pyx @@ -113,10 +113,9 @@ cdef inline tuple _query_peer_access_ids(DeviceMemoryResource mr): cdef inline bint _peer_access_includes(DeviceMemoryResource mr, int dev_id): """Return True if peer access from ``dev_id`` is currently granted.""" cdef cydriver.CUmemAccess_flags flags - cdef cydriver.CUmemLocation location = cydriver.CUmemLocation( - type=cydriver.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE, - id=dev_id, - ) + cdef cydriver.CUmemLocation location + location.type = cydriver.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE + location.id = dev_id with nogil: HANDLE_RETURN(cydriver.cuMemPoolGetAccess(&flags, as_cu(mr._h_pool), &location)) return flags == cydriver.CUmemAccess_flags.CU_MEM_ACCESS_FLAGS_PROT_READWRITE diff --git a/cuda_core/cuda/core/graph/_graph_node.pyx b/cuda_core/cuda/core/graph/_graph_node.pyx index 2c9c07e6b3a..c4b6b02bf37 100644 --- a/cuda_core/cuda/core/graph/_graph_node.pyx +++ b/cuda_core/cuda/core/graph/_graph_node.pyx @@ -827,6 +827,7 @@ cdef inline AllocNode GN_alloc(GraphNode self, size_t size, object device, num_deps = 1 cdef vector[cydriver.CUmemAccessDesc] access_descs + cdef cydriver.CUmemAccessDesc access_desc cdef int peer_id cdef list peer_ids = [] @@ -834,13 +835,10 @@ cdef inline AllocNode GN_alloc(GraphNode self, size_t size, object device, for peer_dev in peer_access: peer_id = getattr(peer_dev, 'device_id', peer_dev) peer_ids.append(peer_id) - access_descs.push_back(cydriver.CUmemAccessDesc_st( - cydriver.CUmemLocation_st( - cydriver.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE, - peer_id - ), - cydriver.CUmemAccess_flags.CU_MEM_ACCESS_FLAGS_PROT_READWRITE - )) + access_desc.location.type = cydriver.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE + access_desc.location.id = peer_id + access_desc.flags = cydriver.CUmemAccess_flags.CU_MEM_ACCESS_FLAGS_PROT_READWRITE + access_descs.push_back(access_desc) cdef str memory_type_str = "device" if memory_type is None else str(memory_type)