Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
19 commits
Select commit Hold shift + click to select a range
2b4c46f
fix(cuda.core): move VirtualMemoryResource onto the _rt handle layer
Andy-Jost Sep 18, 2026
e7cbacc
test(cuda.core): run the VMM shutdown test from an empty directory
Andy-Jost Sep 18, 2026
f9e5940
fix(cuda.core): harden VirtualMemoryResource after review
Andy-Jost Sep 18, 2026
2a1db37
Merge remote-tracking branch 'origin/main' into ajost/vmm-redesign
Andy-Jost Sep 18, 2026
d5b7977
test(cuda.core): pass the handle-type enum to the CUDA 13.0 bindings
Andy-Jost Sep 18, 2026
5f07a73
test(cuda.core): cover host VMM grow, default-stream capture skip, an…
Andy-Jost Sep 21, 2026
6869185
Merge remote-tracking branch 'origin/main' into ajost/vmm-redesign
Andy-Jost Sep 21, 2026
1f9c014
docs(cuda.core): state the VirtualMemoryResource invariants in VMM_DE…
Andy-Jost Sep 21, 2026
7e08822
refactor(cuda.core): make VMM resource attributes read-only and share…
Andy-Jost Sep 28, 2026
343a495
test(cuda.core): check VMM release per address instead of device-wide…
Andy-Jost Sep 28, 2026
312b0b4
fix(cuda.core): make VMM ranges immutable and give each buffer its own
Andy-Jost Sep 29, 2026
686b7be
Merge remote-tracking branch 'origin/main' into ajost/vmm-redesign
Andy-Jost Sep 29, 2026
beea766
fix(cuda.core): return an empty stream for an empty device pointer ha…
Andy-Jost Sep 29, 2026
988e342
test(cuda.core): bound the joins in the concurrent VMM grow test
Andy-Jost Sep 29, 2026
d226120
Merge remote-tracking branch 'origin/main' into ajost/vmm-redesign
Andy-Jost Oct 1, 2026
53479be
fix(cuda.core): keep the pending exception across a MemoryResource de…
Andy-Jost Oct 1, 2026
6cb1826
docs(cuda.core): restore __weakref__ on VirtualMemoryResource and sta…
Andy-Jost Oct 1, 2026
012c0ff
Merge remote-tracking branch 'origin/main' into ajost/vmm-redesign
Andy-Jost Oct 1, 2026
818a440
fix(cuda.core): query capture state with cuStreamIsCapturing in the V…
Andy-Jost Oct 1, 2026
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
5 changes: 5 additions & 0 deletions cuda_core/cuda/core/_cpp/rt/DESIGN.md
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,11 @@ Internally, handles use **shared pointer aliasing**: the actual managed object i
"box" containing the resource, its dependencies, and any state needed for destruction.
The public handle points only to the raw resource field, keeping the API minimal.

The virtual memory resource adds `MemAllocationHandle`, `VaReservationHandle` and
`VaMappingHandle`. Their values are `TaggedHandle<T, N>` wrappers, because
`CUmemGenericAllocationHandle` and `CUdeviceptr` are both `unsigned long long` and the
accessor overloads must stay distinct. See [VMM_DESIGN.md](VMM_DESIGN.md).

### Why shared_ptr?

- **Automatic reference counting**: Resources are released when the last reference
Expand Down
269 changes: 269 additions & 0 deletions cuda_core/cuda/core/_cpp/rt/VMM_DESIGN.md

Large diffs are not rendered by default.

66 changes: 66 additions & 0 deletions cuda_core/cuda/core/_cpp/rt/api.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
#pragma once

#include "types.hpp"
#include <vector>
#include <cuda.h>
#include <nvrtc.h>
#include <cstddef>
Expand Down Expand Up @@ -232,6 +233,7 @@ DevicePtrHandle deviceptr_import_ipc(

// Access the deallocation stream for a device pointer handle (read-only).
// For non-owning handles, the stream is not used but can still be accessed.
// Returns an empty handle for an empty device pointer handle.
StreamHandle deallocation_stream(const DevicePtrHandle& h) noexcept;

// Set the deallocation stream for a device pointer handle.
Expand All @@ -240,6 +242,70 @@ StreamHandle deallocation_stream(const DevicePtrHandle& h) noexcept;
CUresult set_deallocation_stream(
const DevicePtrHandle& h, const StreamHandle& h_stream) noexcept;

// ============================================================================
// Virtual memory management (VMM_DESIGN.md)
//
// A VirtualMemoryResource buffer is a range of mappings. Each mapping holds
// one physical allocation and one address reservation; the mapping deleter
// unmaps, then the allocation is released and the reservation freed as their
// last references go. A buffer's DevicePtrHandle owns the range. Ranges are
// immutable: a grow builds a new range for its result.
// ============================================================================

// Create a physical allocation via cuMemCreate. The access descriptors are
// applied to every mapping of this allocation. When the last reference is
// released, cuMemRelease is called; the memory is freed once no mapping
// remains. Returns empty handle on error (caller must check).
MemAllocationHandle create_mem_allocation_handle(size_t size, const CUmemAllocationProp& prop,
const CUmemAccessDesc* descs, size_t count);

// Size of the allocation; the only size cuMemMap accepts for it.
size_t mem_allocation_size(const MemAllocationHandle& h) noexcept;

// Reserve an address range via cuMemAddressReserve. Pass alignment 0 for the
// driver default. When the last reference is released, cuMemAddressFree is
// called with the exact reserved pair. Returns empty handle on error.
VaReservationHandle create_va_reservation_handle(size_t size, size_t alignment, CUdeviceptr hint);

// Size of the reservation.
size_t va_reservation_size(const VaReservationHandle& h) noexcept;

// Map the whole allocation at ptr inside the reservation via cuMemMap and
// apply the allocation's access descriptors. The mapping structurally depends
// on both handles. When the last reference is released, cuMemUnmap is called
// first. Returns empty handle on error, including a range outside the
// reservation; a failed cuMemSetAccess unmaps before returning.
VaMappingHandle create_va_mapping_handle(CUdeviceptr ptr, const MemAllocationHandle& h_alloc,
const VaReservationHandle& h_res);

// Mapping accessors.
size_t va_mapping_size(const VaMappingHandle& h) noexcept;
MemAllocationHandle va_mapping_allocation(const VaMappingHandle& h) noexcept;

// Build an immutable range from mappings in ascending, contiguous order. A
// grow builds a new range for its result and never changes the input's.
// May throw std::bad_alloc.
VmmRangeHandle create_vmm_range(const std::vector<VaMappingHandle>& mappings);

// The range of a device pointer handle created by deviceptr_create_vmm. Only
// for such handles: the caller (VirtualMemoryBuffer) guarantees the origin.
// Empty for an empty handle.
VmmRangeHandle vmm_range(const DevicePtrHandle& h) noexcept;

// Range accessors: a copy of the mapping list, which a grow extends and turns
// into a new range, and the range total. Reads of an immutable range need no
// synchronization.
std::vector<VaMappingHandle> vmm_range_mappings(const VmmRangeHandle& range); // may throw
size_t vmm_range_total(const VmmRangeHandle& range) noexcept;

// Create a device pointer handle whose box owns a range (an empty range for
// a size-zero buffer). The box records no deallocation stream; set one with
// set_deallocation_stream. When the last reference is released, the recorded
// stream is synchronized (skipped with a report when that would disturb a
// capture) and the box freed, which unmaps every mapping this buffer was the
// last to hold. Returns empty handle for a null range handle.
DevicePtrHandle deviceptr_create_vmm(CUdeviceptr base, const VmmRangeHandle& range);

// ============================================================================
// Library handle functions
// ============================================================================
Expand Down
14 changes: 13 additions & 1 deletion cuda_core/cuda/core/_cpp/rt/driver_api.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -126,7 +126,19 @@ namespace cuda_core::rt {
X(cuTexObjectCreate, 5000) \
X(cuTexObjectDestroy, 5000) \
X(cuSurfObjectCreate, 5000) \
X(cuSurfObjectDestroy, 5000)
X(cuSurfObjectDestroy, 5000) \
/* Virtual memory management (VMM_DESIGN.md) */ \
X(cuMemCreate, 10020) \
X(cuMemRelease, 10020) \
X(cuMemAddressReserve, 10020) \
X(cuMemAddressFree, 10020) \
X(cuMemMap, 10020) \
X(cuMemUnmap, 10020) \
X(cuMemSetAccess, 10020) \
/* cuda-bindings requests 7000 (PTDS) or 2000 (legacy) */ \
X(cuStreamSynchronize, 7000) \
X(cuStreamIsCapturing, 10000) \
X(cuThreadExchangeStreamCaptureMode, 10010)

#define CUDA_CORE_DECLARE_DRIVER_FN(name, introduced) extern decltype(&name) p_##name;
CUDA_CORE_DRIVER_FUNCTIONS(CUDA_CORE_DECLARE_DRIVER_FN)
Expand Down
9 changes: 9 additions & 0 deletions cuda_core/cuda/core/_cpp/rt/internal.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,12 @@ ContextHandle deallocation_context(const DeallocationStream& stream) noexcept;
// Implemented in stream.cpp
bool make_deallocation_stream(const StreamHandle& h, DeallocationStream& out) noexcept;

// Implemented in virtual_memory.cpp. Synchronize the stream a VMM buffer
// recorded before its mappings are released. Skips the sync, with a report
// when a capture is the reason, if no stream was recorded, the interpreter is
// finalizing, or the sync would disturb a graph capture (VMM_DESIGN.md).
void vmm_sync_before_release(const DeallocationStream& stream) noexcept;

// Decorate a status-returning cleanup call to report whenever it fails. CUDA
// calls (CUresult) are reported with the error name and description; NVRTC,
// NVVM and nvJitLink calls (integer status codes) with the raw code.
Expand Down Expand Up @@ -126,6 +132,9 @@ const WarnOnFailure<p_cuSurfObjectDestroy> pw_cuSurfObjectDestroy{"cuSurfObjectD
const WarnOnFailure<p_cuGreenCtxDestroy> pw_cuGreenCtxDestroy{"cuGreenCtxDestroy"};
const WarnOnFailure<p_cuMemPoolDestroy> pw_cuMemPoolDestroy{"cuMemPoolDestroy"};
const WarnOnFailure<p_cuMemFreeHost> pw_cuMemFreeHost{"cuMemFreeHost"};
const WarnOnFailure<p_cuMemRelease> pw_cuMemRelease{"cuMemRelease"};
const WarnOnFailure<p_cuMemUnmap> pw_cuMemUnmap{"cuMemUnmap"};
const WarnOnFailure<p_cuMemAddressFree> pw_cuMemAddressFree{"cuMemAddressFree"};
const WarnOnFailure<p_cuGraphDestroy> pw_cuGraphDestroy{"cuGraphDestroy"};
const WarnOnFailure<p_cuGraphExecDestroy> pw_cuGraphExecDestroy{"cuGraphExecDestroy"};
const WarnOnFailure<p_cuGraphicsUnregisterResource> pw_cuGraphicsUnregisterResource{"cuGraphicsUnregisterResource"};
Expand Down
48 changes: 46 additions & 2 deletions cuda_core/cuda/core/_cpp/rt/memory.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -123,6 +123,15 @@ struct DevicePtrBox {
// stream tokens carry a bound context.
mutable DeallocationStream deallocation;
};

// The box behind a VirtualMemoryResource buffer: the range of mappings it
// owns (VMM_DESIGN.md). Every handle on a VirtualMemoryBuffer comes from
// deviceptr_create_vmm, so vmm_range() downcasts without a tag; there are no
// virtual functions, so other boxes pay nothing. A size-zero buffer has an
// empty range.
struct VmmDevicePtrBox : DevicePtrBox {
VmmRangeHandle range;
};
} // namespace

// Recovers the owning DevicePtrBox from the aliased CUdeviceptr pointer.
Expand All @@ -137,9 +146,10 @@ static DevicePtrBox* get_box(const DevicePtrHandle& h) {
);
}

// Return the stream that orders a device pointer's deallocation.
// Return the stream that orders a device pointer's deallocation; empty for an
// empty handle, as set_deallocation_stream rejects one.
StreamHandle deallocation_stream(const DevicePtrHandle& h) noexcept {
return get_box(h)->deallocation.h_stream;
return h ? get_box(h)->deallocation.h_stream : StreamHandle{};
}

// Replace the stream that orders a device pointer's deallocation.
Expand Down Expand Up @@ -300,6 +310,36 @@ DevicePtrHandle deviceptr_create_mapped_graphics(
return DevicePtrHandle(box, &box->resource);
}

// ============================================================================
// Virtual memory ranges (VMM_DESIGN.md)
// ============================================================================

DevicePtrHandle deviceptr_create_vmm(CUdeviceptr base, const VmmRangeHandle& range) {
if (!range) { // a null handle; an empty range is valid
err = CUDA_ERROR_INVALID_VALUE;
return {};
}
auto box = std::shared_ptr<VmmDevicePtrBox>(
new VmmDevicePtrBox{{base, DeallocationStream{}}, range},
[](VmmDevicePtrBox* b) {
GILReleaseGuard gil;
// Order the release on this buffer's recorded stream, then free
// the box, which drops the range: every mapping this buffer was
// the last to hold unmaps, and its reservation and allocation
// follow.
vmm_sync_before_release(b->deallocation);
delete b;
}
);
return DevicePtrHandle(box, &box->resource);
}

VmmRangeHandle vmm_range(const DevicePtrHandle& h) noexcept {
// Only for handles from deviceptr_create_vmm; the VirtualMemoryBuffer
// class guarantees that for its callers.
return h ? static_cast<VmmDevicePtrBox*>(get_box(h))->range : VmmRangeHandle{};
}

// ============================================================================
// MemoryResource-owned Device Pointer Handles
// ============================================================================
Expand All @@ -325,6 +365,10 @@ DevicePtrHandle deviceptr_create_with_mr(CUdeviceptr ptr, size_t size, PyObject*
[mr, size](DevicePtrBox* b) {
GILAcquireGuard gil;
if (gil.acquired()) {
// The last reference may go while an exception propagates
// through the releasing caller; deallocate() must run with a
// clean error state and leave that exception in place.
PendingExceptionGuard pending;
if (mr_dealloc_cb) {
const DeallocationStream& stream = b->deallocation;
cleanup_in_context(
Expand Down
51 changes: 51 additions & 0 deletions cuda_core/cuda/core/_cpp/rt/py.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,45 @@ class GILAcquireGuard {
bool acquired_;
};

// Save the Python exception in flight and restore it on scope exit, dropping
// anything the scope itself raised. A deleter may run while an exception is
// propagating through the caller that released the last reference; Python
// code it calls (a MemoryResource's deallocate, the warnings machinery) must
// start with a clean error state and must not leave that exception cleared,
// or the caller returns an error with no exception set. Allocates nothing.
// Construct with the GIL held and destroy before releasing it.
class PendingExceptionGuard {
public:
PendingExceptionGuard() noexcept {
#if PY_VERSION_HEX >= 0x030C0000
pending_ = PyErr_GetRaisedException();
#else
PyErr_Fetch(&type_, &value_, &tb_);
#endif
}

~PendingExceptionGuard() {
PyErr_Clear();
#if PY_VERSION_HEX >= 0x030C0000
PyErr_SetRaisedException(pending_);
#else
PyErr_Restore(type_, value_, tb_);
#endif
}

PendingExceptionGuard(const PendingExceptionGuard&) = delete;
PendingExceptionGuard& operator=(const PendingExceptionGuard&) = delete;

private:
#if PY_VERSION_HEX >= 0x030C0000
PyObject* pending_ = nullptr;
#else
PyObject* type_ = nullptr;
PyObject* value_ = nullptr;
PyObject* tb_ = nullptr;
#endif
};

// as_py() - convert handle to Python wrapper object (returns new reference)
namespace detail {
// n.b. class lookup is not cached to avoid deadlock hazard, see DESIGN.md
Expand Down Expand Up @@ -205,6 +244,18 @@ inline PyObject* as_py(const SurfObjectHandle& h) noexcept {
return detail::make_py("cuda.bindings.driver", "CUsurfObject", as_intptr(h));
}

inline PyObject* as_py(const MemAllocationHandle& h) noexcept {
return detail::make_py("cuda.bindings.driver", "CUmemGenericAllocationHandle", as_intptr(h));
}

inline PyObject* as_py(const VaReservationHandle& h) noexcept {
return detail::make_py("cuda.bindings.driver", "CUdeviceptr", as_intptr(h));
}

inline PyObject* as_py(const VaMappingHandle& h) noexcept {
return detail::make_py("cuda.bindings.driver", "CUdeviceptr", as_intptr(h));
}

// ============================================================================
// Python-coupled API: the prototypes that take or return PyObject*
// ============================================================================
Expand Down
33 changes: 0 additions & 33 deletions cuda_core/cuda/core/_cpp/rt/py_driver_fns.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -48,39 +48,6 @@ std::atomic<bool> unavailable_reported[kTables];
std::mutex fill_mutex;
char fill_error[kTables][512] = {}; // guarded by fill_mutex

// Saves the pending Python exception on construction and restores it on
// destruction. The Python calls in between start from a clean error state, and
// the caller's exception survives. Requires the GIL.
class PendingExceptionGuard {
public:
PendingExceptionGuard() noexcept {
#if PY_VERSION_HEX >= 0x030C0000
exc_ = PyErr_GetRaisedException();
#else
PyErr_Fetch(&type_, &value_, &traceback_);
#endif
}
~PendingExceptionGuard() {
PyErr_Clear(); // drop anything the guarded calls left set
#if PY_VERSION_HEX >= 0x030C0000
PyErr_SetRaisedException(exc_);
#else
PyErr_Restore(type_, value_, traceback_);
#endif
}
PendingExceptionGuard(const PendingExceptionGuard&) = delete;
PendingExceptionGuard& operator=(const PendingExceptionGuard&) = delete;

private:
#if PY_VERSION_HEX >= 0x030C0000
PyObject* exc_ = nullptr;
#else
PyObject* type_ = nullptr;
PyObject* value_ = nullptr;
PyObject* traceback_ = nullptr;
#endif
};

std::size_t index_of(FnTable table) noexcept { return static_cast<std::size_t>(table); }

const char* module_name(FnTable table) noexcept {
Expand Down
12 changes: 1 addition & 11 deletions cuda_core/cuda/core/_cpp/rt/py_report.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -34,12 +34,7 @@ void report_message(const char* message) noexcept {
GILAcquireGuard gil;
if (gil.acquired()) {
// Deleters can run while a Python exception is propagating; keep it.
#if PY_VERSION_HEX >= 0x030C0000
PyObject* pending = PyErr_GetRaisedException();
#else
PyObject *pending_type, *pending_value, *pending_tb;
PyErr_Fetch(&pending_type, &pending_value, &pending_tb);
#endif
PendingExceptionGuard pending;
bool interrupted = false;
if (PyErr_WarnEx(category, message, 1) != 0) {
interrupted = PyErr_ExceptionMatches(PyExc_KeyboardInterrupt);
Expand All @@ -52,11 +47,6 @@ void report_message(const char* message) noexcept {
Py_XDECREF(subject);
}
}
#if PY_VERSION_HEX >= 0x030C0000
PyErr_SetRaisedException(pending);
#else
PyErr_Restore(pending_type, pending_value, pending_tb);
#endif
if (interrupted) {
PyErr_SetInterrupt();
}
Expand Down
Loading
Loading