Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
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
14 changes: 14 additions & 0 deletions cuda-core/src/simt/context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -548,6 +548,20 @@ impl CudaContext {
}
}

/// Queries the device memory of this context as `(free, total)` bytes.
///
/// Binds the context first, then calls `cuMemGetInfo`. `total` is the
/// device's memory; `free` is what the driver can hand out at this
/// moment, so allocations on other streams and driver-side caches move
/// it between two calls.
pub fn mem_info(&self) -> Result<(usize, usize), DriverError> {
self.bind_to_thread()?;
let mut free = 0usize;
let mut total = 0usize;
unsafe { cuda_bindings::cuMemGetInfo_v2(&mut free, &mut total) }.result()?;
Ok((free, total))
}

/// Queries dimension, thread-count, and portable shared-memory launch
/// limits for this device.
///
Expand Down
29 changes: 21 additions & 8 deletions cuda-core/tests/simt_device_buffer_leaks.rs
Original file line number Diff line number Diff line change
Expand Up @@ -67,19 +67,13 @@ fn ctx_strong_count_returns_to_baseline_after_buffer_lifecycle() {
/// This cannot deterministically force the async-enqueue failure that
/// triggered the original leak (no public API constructs an invalid
/// `CudaStream`), but it pins the happy-path accounting with
/// `cuMemGetInfo` before and after the cycles.
/// [`CudaContext::mem_info`] before and after the cycles.
#[test]
fn vram_returns_to_baseline_after_buffer_cycles() {
let ctx = CudaContext::new(0).expect("failed to create CUDA context");
let stream = ctx.new_stream().expect("failed to create CUDA stream");

fn free_mem() -> usize {
let mut free = 0usize;
let mut total = 0usize;
let rc = unsafe { cuda_bindings::cuMemGetInfo_v2(&mut free, &mut total) };
assert_eq!(rc, 0, "cuMemGetInfo failed: {rc}");
free
}
let free_mem = || ctx.mem_info().expect("mem_info failed").0;

// Warm up driver allocator caches so the measured window is stable.
for _ in 0..4 {
Expand Down Expand Up @@ -135,3 +129,22 @@ fn zero_length_construction_succeeds_for_both_constructors() {
"empty buffers must not leak a ctx strong count"
);
}

/// `mem_info` reports the device's memory: `total` matches the device's
/// total memory attribute and `free` never exceeds it.
#[test]
fn mem_info_reports_free_within_the_device_total() {
let ctx = CudaContext::new(0).expect("failed to create CUDA context");
let (free, total) = ctx.mem_info().expect("mem_info failed");

let mut device_total = 0usize;
let rc = unsafe { cuda_bindings::cuDeviceTotalMem_v2(&mut device_total, ctx.cu_device()) };
assert_eq!(rc, 0, "cuDeviceTotalMem failed: {rc}");

assert_eq!(total, device_total, "total must be the device's memory");
assert!(
free <= total,
"free ({free}) must not exceed total ({total})"
);
assert!(total > 0, "a device has memory");
}
5 changes: 5 additions & 0 deletions cutile-rs/CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,11 @@ and this project adheres to [Semantic Versioning](https://semver.org/).

## [Unreleased]

### Added

- `CudaContext::mem_info`, the `(free, total)` device memory query
(`cuMemGetInfo`) for the context.

### Changed

- cuTile Rust now lives in the [NVIDIA/cuda-rust](https://github.com/NVIDIA/cuda-rust)
Expand Down