diff --git a/cuda_core/cuda/core/graph/_subclasses.pyi b/cuda_core/cuda/core/graph/_subclasses.pyi index e380307b670..41c70463a29 100644 --- a/cuda_core/cuda/core/graph/_subclasses.pyi +++ b/cuda_core/cuda/core/graph/_subclasses.pyi @@ -187,7 +187,7 @@ class MemcpyNode(GraphNode): Omitted parameters preserve their current values. ``dst_owner`` and ``src_owner`` may only accompany their corresponding raw addresses. Multidimensional, pitched, offset, and array-backed memcpy nodes are - not supported. + not supported. Nodes recorded by stream capture are supported. With drivers from CUDA 12.2 through 13.1, the node's intended CUDA context must be current when this method runs. With the CUDA 13 build @@ -397,7 +397,10 @@ class ExecutableMemsetNode(ExecutableGraphNode): class ExecutableMemcpyNode(ExecutableGraphNode): """An executable memcpy-node view.""" def update(self, *, dst: Buffer | int, src: Buffer | int, size: int) -> None: - """Replace all one-dimensional memcpy parameters for future launches.""" + """Replace all one-dimensional memcpy parameters for future launches. + + Nodes recorded by stream capture are supported. + """ @property def is_enabled(self) -> bool: """Whether this node is enabled in the executable graph.""" diff --git a/cuda_core/cuda/core/graph/_subclasses.pyx b/cuda_core/cuda/core/graph/_subclasses.pyx index 540849ee97a..54df8c74d2e 100644 --- a/cuda_core/cuda/core/graph/_subclasses.pyx +++ b/cuda_core/cuda/core/graph/_subclasses.pyx @@ -243,13 +243,33 @@ cdef void _reject_unsupported_kernel_node( "updating clustered or cooperative kernel nodes is not supported") +cdef inline bint _is_linear_memcpy_memory_type( + cydriver.CUmemorytype memory_type) noexcept nogil: + # HOST and DEVICE name the location of a linear operand. UNIFIED is what + # the driver records for a copy captured from cuMemcpyAsync: a generic + # address that it resolves itself, stored in the device-pointer field. + return (memory_type == cydriver.CU_MEMORYTYPE_HOST + or memory_type == cydriver.CU_MEMORYTYPE_DEVICE + or memory_type == cydriver.CU_MEMORYTYPE_UNIFIED) + + +cdef str _memcpy_memory_type_tag(cydriver.CUmemorytype memory_type): + if memory_type == cydriver.CU_MEMORYTYPE_HOST: + return "H" + if memory_type == cydriver.CU_MEMORYTYPE_DEVICE: + return "D" + if memory_type == cydriver.CU_MEMORYTYPE_UNIFIED: + return "U" + if memory_type == cydriver.CU_MEMORYTYPE_ARRAY: + return "A" + return str(int(memory_type)) + + cdef bint _is_supported_memcpy_descriptor( cydriver.CUDA_MEMCPY3D* params) noexcept nogil: return ( - (params.srcMemoryType == cydriver.CU_MEMORYTYPE_HOST or - params.srcMemoryType == cydriver.CU_MEMORYTYPE_DEVICE) - and (params.dstMemoryType == cydriver.CU_MEMORYTYPE_HOST or - params.dstMemoryType == cydriver.CU_MEMORYTYPE_DEVICE) + _is_linear_memcpy_memory_type(params.srcMemoryType) + and _is_linear_memcpy_memory_type(params.dstMemoryType) and params.srcXInBytes == 0 and params.srcY == 0 and params.srcZ == 0 @@ -819,8 +839,8 @@ cdef class MemcpyNode(GraphNode): params.dstMemoryType, params.srcMemoryType) def __repr__(self) -> str: - cdef str dt = "H" if self._dst_type == cydriver.CU_MEMORYTYPE_HOST else "D" - cdef str st = "H" if self._src_type == cydriver.CU_MEMORYTYPE_HOST else "D" + cdef str dt = _memcpy_memory_type_tag(self._dst_type) + cdef str st = _memcpy_memory_type_tag(self._src_type) return (f"") @@ -838,7 +858,7 @@ cdef class MemcpyNode(GraphNode): Omitted parameters preserve their current values. ``dst_owner`` and ``src_owner`` may only accompany their corresponding raw addresses. Multidimensional, pitched, offset, and array-backed memcpy nodes are - not supported. + not supported. Nodes recorded by stream capture are supported. With drivers from CUDA 12.2 through 13.1, the node's intended CUDA context must be current when this method runs. With the CUDA 13 build @@ -894,32 +914,31 @@ cdef class MemcpyNode(GraphNode): "updating multidimensional, pitched, offset, or array-backed " "memcpy nodes is not supported") + # The descriptor check above admits HOST, DEVICE, and UNIFIED operands; + # the last two keep their address in the device-pointer field. c_dst_type = params.memcpy.copyParams.dstMemoryType c_src_type = params.memcpy.copyParams.srcMemoryType if c_dst_type == cydriver.CU_MEMORYTYPE_HOST: c_dst = ( params.memcpy.copyParams.dstHost) - elif c_dst_type == cydriver.CU_MEMORYTYPE_DEVICE: - c_dst = params.memcpy.copyParams.dstDevice else: - raise NotImplementedError( - f"unsupported destination memory type: {int(c_dst_type)}") + c_dst = params.memcpy.copyParams.dstDevice if c_src_type == cydriver.CU_MEMORYTYPE_HOST: c_src = ( params.memcpy.copyParams.srcHost) - elif c_src_type == cydriver.CU_MEMORYTYPE_DEVICE: - c_src = params.memcpy.copyParams.srcDevice else: - raise NotImplementedError( - f"unsupported source memory type: {int(c_src_type)}") + c_src = params.memcpy.copyParams.srcDevice HANDLE_RETURN(graph_get_attachment( h_graph, node, &dst_attachment_owner, &src_attachment_owner)) + # A unified operand stays unified: any address is valid for that type, + # and an executable update rejects a change of memory type. if dst is not None: dst_attachment_owner = _resolve_memcpy_operand( dst, dst_owner, "dst", &c_dst) - c_dst_type = _get_memcpy_memory_type(c_dst) + if c_dst_type != cydriver.CU_MEMORYTYPE_UNIFIED: + c_dst_type = _get_memcpy_memory_type(c_dst) params.memcpy.copyParams.dstMemoryType = c_dst_type params.memcpy.copyParams.dstHost = NULL params.memcpy.copyParams.dstDevice = 0 @@ -932,7 +951,8 @@ cdef class MemcpyNode(GraphNode): if src is not None: src_attachment_owner = _resolve_memcpy_operand( src, src_owner, "src", &c_src) - c_src_type = _get_memcpy_memory_type(c_src) + if c_src_type != cydriver.CU_MEMORYTYPE_UNIFIED: + c_src_type = _get_memcpy_memory_type(c_src) params.memcpy.copyParams.srcMemoryType = c_src_type params.memcpy.copyParams.srcHost = NULL params.memcpy.copyParams.srcDevice = 0 @@ -1551,7 +1571,10 @@ cdef class ExecutableMemcpyNode(ExecutableGraphNode): src: Buffer | int, size_t size, ) -> None: - """Replace all one-dimensional memcpy parameters for future launches.""" + """Replace all one-dimensional memcpy parameters for future launches. + + Nodes recorded by stream capture are supported. + """ cdef cydriver.CUdeviceptr c_dst cdef cydriver.CUdeviceptr c_src cdef OpaqueHandle dst_owner = _resolve_memcpy_operand( @@ -1562,12 +1585,29 @@ cdef class ExecutableMemcpyNode(ExecutableGraphNode): cdef cydriver.CUmemorytype src_type cdef cydriver.CUcontext ctx = NULL cdef cydriver.CUgraphNodeParams params + cdef cydriver.CUDA_MEMCPY3D recorded + cdef cydriver.CUgraphNode node = as_cu(self._h_node) c_memset(¶ms, 0, sizeof(params)) params.type = cydriver.CU_GRAPH_NODE_TYPE_MEMCPY _init_memcpy_params( c_dst, c_src, size, ¶ms.memcpy.copyParams, &dst_type, &src_type) + # A destroyed node is reported by _set_executable_node_params below. + if node != NULL: + with nogil: + HANDLE_RETURN(cydriver.cuGraphMemcpyNodeGetParams( + node, &recorded)) + if recorded.dstMemoryType == cydriver.CU_MEMORYTYPE_UNIFIED: + params.memcpy.copyParams.dstMemoryType = ( + cydriver.CU_MEMORYTYPE_UNIFIED) + params.memcpy.copyParams.dstHost = NULL + params.memcpy.copyParams.dstDevice = c_dst + if recorded.srcMemoryType == cydriver.CU_MEMORYTYPE_UNIFIED: + params.memcpy.copyParams.srcMemoryType = ( + cydriver.CU_MEMORYTYPE_UNIFIED) + params.memcpy.copyParams.srcHost = NULL + params.memcpy.copyParams.srcDevice = c_src with nogil: HANDLE_RETURN(cydriver.cuCtxGetCurrent(&ctx)) params.memcpy.copyCtx = ctx diff --git a/cuda_core/docs/source/release/1.3.0-notes.rst b/cuda_core/docs/source/release/1.3.0-notes.rst index 81321e71006..5e3cbd0baa6 100644 --- a/cuda_core/docs/source/release/1.3.0-notes.rst +++ b/cuda_core/docs/source/release/1.3.0-notes.rst @@ -257,3 +257,13 @@ Fixes and enhancements - The ``StridedMemoryView`` DLPack C exchange API now preserves the underlying Python exception when an export or import fails, as required by the DLPack contract, instead of returning ``-1`` with no exception set. + +- :meth:`~graph.MemcpyNode.update` and the executable view of a memcpy node + now accept nodes recorded by stream capture. The driver stores the operands + of a captured :meth:`Buffer.copy_from` or :meth:`Buffer.copy_to` as unified + addresses, and ``update`` rejected that descriptor with a message about + multidimensional copies. Such an operand stays unified when it is replaced, + because the driver does not accept a change of an operand's memory type in + an instantiated graph, so :meth:`~graph.Graph.update` applies the change as + well. The node's ``repr`` now shows ``U`` for unified operands. + (`#2649 `__) diff --git a/cuda_core/tests/graph/test_graph_node_update.py b/cuda_core/tests/graph/test_graph_node_update.py index 8e0653f67be..0114be8e96c 100644 --- a/cuda_core/tests/graph/test_graph_node_update.py +++ b/cuda_core/tests/graph/test_graph_node_update.py @@ -745,6 +745,139 @@ def test_memcpy_update_between_host_and_device(init_cuda, device_operand): assert list(host_dst_bytes) == [0x5A] * 4 +def _capture_device_memcpy(device, stream, size=64): + """Capture ``dst.copy_from(src)`` between two device buffers. + + The driver records both operands of a captured ``cuMemcpyAsync`` as + unified addresses rather than as device memory (#2649). The builder is + returned so that the captured graph stays alive with the definition view. + """ + src = device.memory_resource.allocate(size, stream=stream) + dst = device.memory_resource.allocate(size, stream=stream) + src.fill(0x5A, stream=stream) + dst.fill(0, stream=stream) + stream.sync() + + builder = device.create_graph_builder().begin_building() + dst.copy_from(src, stream=builder) + builder.end_building() + graph_def = builder.graph_definition + node = next(n for n in graph_def.nodes() if isinstance(n, MemcpyNode)) + return builder, graph_def, node, dst, src + + +def _read_device_bytes(buffer, stream): + host = LegacyPinnedMemoryResource().allocate(buffer.size) + buffer.copy_to(host, stream=stream) + stream.sync() + return list((ctypes.c_uint8 * buffer.size).from_address(int(host.handle))) + + +_MEMORY_TYPE_TAGS = { + driver.CUmemorytype.CU_MEMORYTYPE_HOST: "H", + driver.CUmemorytype.CU_MEMORYTYPE_DEVICE: "D", + driver.CUmemorytype.CU_MEMORYTYPE_UNIFIED: "U", +} + + +@pytest.mark.agent_authored(model="claude-fable-5-1") +def test_memcpy_update_captured_node_size(init_cuda): + """A memcpy node recorded by stream capture accepts a size update (#2649).""" + if driver_version() < (12, 2, 0): + pytest.skip("individual graph node updates require CUDA 12.2+") + + stream = init_cuda.create_stream() + builder, graph_def, node, dst, src = _capture_device_memcpy(init_cuda, stream) + recorded = handle_return(driver.cuGraphMemcpyNodeGetParams(node.handle)) + assert node.size == 64 + assert node.dst == int(dst.handle) + assert node.src == int(src.handle) + assert f"({_MEMORY_TYPE_TAGS[recorded.dstMemoryType]})" in repr(node) + assert f"({_MEMORY_TYPE_TAGS[recorded.srcMemoryType]})" in repr(node) + + node.update(size=32) + + assert node.size == 32 + assert node.dst == int(dst.handle) + assert node.src == int(src.handle) + updated = handle_return(driver.cuGraphMemcpyNodeGetParams(node.handle)) + assert updated.WidthInBytes == 32 + assert updated.dstMemoryType == recorded.dstMemoryType + assert updated.srcMemoryType == recorded.srcMemoryType + + graph = graph_def.instantiate() + graph.launch(stream) + assert _read_device_bytes(dst, stream) == [0x5A] * 32 + [0] * 32 + + +@pytest.mark.parametrize("operand", ["src", "dst"]) +@pytest.mark.agent_authored(model="claude-fable-5-1") +def test_memcpy_update_captured_node_operand(init_cuda, operand): + """Replacing one operand of a captured memcpy node keeps the other as recorded (#2649).""" + if driver_version() < (12, 2, 0): + pytest.skip("individual graph node updates require CUDA 12.2+") + + stream = init_cuda.create_stream() + builder, graph_def, node, dst, src = _capture_device_memcpy(init_cuda, stream) + recorded = handle_return(driver.cuGraphMemcpyNodeGetParams(node.handle)) + instantiated_before = graph_def.instantiate() + replacement = init_cuda.memory_resource.allocate(64, stream=stream) + replacement.fill(0xA5 if operand == "src" else 0, stream=stream) + stream.sync() + + if operand == "src": + node.update(src=replacement) + assert node.src == int(replacement.handle) + assert node.dst == int(dst.handle) + copied_into, expected = dst, [0xA5] * 64 + else: + node.update(dst=replacement) + assert node.dst == int(replacement.handle) + assert node.src == int(src.handle) + copied_into, expected = replacement, [0x5A] * 64 + assert node.size == 64 + + # The driver rejects a change of memory type in an executable update, so + # the replaced operand keeps the recorded type. + updated = handle_return(driver.cuGraphMemcpyNodeGetParams(node.handle)) + assert updated.dstMemoryType == recorded.dstMemoryType + assert updated.srcMemoryType == recorded.srcMemoryType + + fresh = graph_def.instantiate() + fresh.launch(stream) + assert _read_device_bytes(copied_into, stream) == expected + + copied_into.fill(0, stream=stream) + stream.sync() + instantiated_before.update(graph_def) + instantiated_before.launch(stream) + assert _read_device_bytes(copied_into, stream) == expected + if operand == "dst": + assert _read_device_bytes(dst, stream) == [0] * 64 + + +@pytest.mark.agent_authored(model="claude-fable-5-1") +def test_executable_memcpy_update_on_captured_node(init_cuda): + """The executable view of a captured memcpy node accepts new operands (#2649).""" + if driver_version() < (12, 2, 0): + pytest.skip("individual graph node updates require CUDA 12.2+") + + stream = init_cuda.create_stream() + builder, graph_def, node, dst, src = _capture_device_memcpy(init_cuda, stream) + new_src = init_cuda.memory_resource.allocate(64, stream=stream) + new_dst = init_cuda.memory_resource.allocate(64, stream=stream) + new_src.fill(0xA5, stream=stream) + new_dst.fill(0, stream=stream) + stream.sync() + + graph = graph_def.instantiate() + graph[node].update(dst=new_dst, src=new_src, size=32) + graph.launch(stream) + + assert _read_device_bytes(new_dst, stream) == [0xA5] * 32 + [0] * 32 + assert _read_device_bytes(dst, stream) == [0] * 64 + + @pytest.mark.agent_authored(model="gpt-5.6") def test_definition_node_update_changes_future_instantiations( definition_update_case,