use std::{
collections::BTreeMap,
ffi::c_void,
ops::{Deref, DerefMut},
sync::{mpsc, Mutex, OnceLock},
};
use crate::{
common::{get_device, set_device_by_id},
error::CudaError,
stream::{device_synchronize, CudaEvent, CudaStream},
};
#[link(name = "cudart")]
extern "C" {
fn cudaHostRegister(ptr: *mut c_void, size: usize, flags: u32) -> i32;
fn cudaHostUnregister(ptr: *mut c_void) -> i32;
}
const CUDA_HOST_REGISTER_PORTABLE: u32 = 0x1;
pub fn register_region(ptr: *mut u8, len: usize) -> bool {
let rc = unsafe { cudaHostRegister(ptr as *mut c_void, len, CUDA_HOST_REGISTER_PORTABLE) };
if rc != 0 {
tracing::debug!(
"cudaHostRegister failed: {}; buffer stays pageable",
CudaError::new(rc)
);
return false;
}
true
}
pub fn unregister_region(ptr: *mut u8) {
unsafe { cudaHostUnregister(ptr as *mut c_void) };
}
pub struct PinnedBuffer {
data: Vec<u8>,
dirty_len: usize,
registered: bool,
last_use: Vec<CudaEvent>,
}
impl PinnedBuffer {
pub fn is_pinned(&self) -> bool {
self.registered
}
pub fn set_dirty_len(&mut self, len: usize) {
self.dirty_len = len;
}
pub fn record_last_use(&mut self, stream: &CudaStream) -> Result<(), CudaError> {
let event = CudaEvent::new()?;
event.record_on(stream)?;
self.last_use.push(event);
Ok(())
}
}
impl Deref for PinnedBuffer {
type Target = [u8];
fn deref(&self) -> &[u8] {
&self.data
}
}
impl DerefMut for PinnedBuffer {
fn deref_mut(&mut self) -> &mut [u8] {
&mut self.data
}
}
impl Drop for PinnedBuffer {
fn drop(&mut self) {
let data = std::mem::take(&mut self.data);
if data.is_empty() {
return;
}
let returned = Returned {
data,
dirty_len: self.dirty_len,
registered: self.registered,
last_use: std::mem::take(&mut self.last_use),
device: get_device().unwrap_or(0),
};
send_to_cleaner(returned);
}
}
fn wait_and_release(returned: Returned) {
let result = if returned.last_use.is_empty() {
set_device_by_id(returned.device).and_then(|_| device_synchronize())
} else {
returned
.last_use
.iter()
.try_for_each(CudaEvent::synchronize)
};
if let Err(e) = result {
tracing::debug!("draining copies from returned buffer failed: {e}");
}
release(returned);
}
fn release(mut returned: Returned) {
if returned.registered {
unregister_region(returned.data.as_mut_ptr());
}
}
struct Returned {
data: Vec<u8>,
dirty_len: usize,
registered: bool,
last_use: Vec<CudaEvent>,
device: i32,
}
const DEFAULT_MAX_POOLED_BYTES: usize = 4 << 30;
struct Pool {
by_size: BTreeMap<usize, Vec<Vec<u8>>>,
total_bytes: usize,
max_bytes: usize,
}
fn pool() -> &'static Mutex<Pool> {
static POOL: OnceLock<Mutex<Pool>> = OnceLock::new();
POOL.get_or_init(|| {
Mutex::new(Pool {
by_size: BTreeMap::new(),
total_bytes: 0,
max_bytes: DEFAULT_MAX_POOLED_BYTES,
})
})
}
pub fn set_max_pooled_bytes(max_bytes: usize) {
let evicted = {
let mut pool = pool().lock().unwrap();
pool.max_bytes = max_bytes;
evict_to_fit(&mut pool, 0)
};
for mut buf in evicted {
unregister_region(buf.as_mut_ptr());
}
}
fn evict_to_fit(pool: &mut Pool, incoming: usize) -> Vec<Vec<u8>> {
let mut evicted = Vec::new();
if incoming > pool.max_bytes {
return evicted;
}
while pool.total_bytes + incoming > pool.max_bytes {
let Some((&class, bufs)) = pool.by_size.iter_mut().next_back() else {
break;
};
evicted.push(bufs.pop().expect("empty size classes are removed"));
if bufs.is_empty() {
pool.by_size.remove(&class);
}
pool.total_bytes -= class;
}
evicted
}
struct Cleaner {
tx: mpsc::Sender<Returned>,
thread: std::thread::JoinHandle<()>,
}
fn cleaner_slot() -> &'static Mutex<Option<Cleaner>> {
static SLOT: OnceLock<Mutex<Option<Cleaner>>> = OnceLock::new();
SLOT.get_or_init(|| Mutex::new(None))
}
fn send_to_cleaner(returned: Returned) {
let mut slot = cleaner_slot().lock().unwrap();
let cleaner = slot.get_or_insert_with(spawn_cleaner);
if let Err(mpsc::SendError(returned)) = cleaner.tx.send(returned) {
tracing::debug!("pinned-cleaner thread is gone; releasing buffer inline");
wait_and_release(returned);
}
}
fn spawn_cleaner() -> Cleaner {
let (tx, rx) = mpsc::channel::<Returned>();
let thread = std::thread::Builder::new()
.name("pinned-cleaner".into())
.spawn(move || {
while let Ok(first) = rx.recv() {
let mut batch = vec![first];
while batch.len() < 64 {
match rx.recv_timeout(std::time::Duration::from_millis(100)) {
Ok(next) => batch.push(next),
Err(_) => break,
}
}
let mut synced_devices = BTreeMap::new();
for returned in batch {
let drained = if !returned.last_use.is_empty() {
returned
.last_use
.iter()
.try_for_each(CudaEvent::synchronize)
.map_err(|e| tracing::debug!("cudaEventSynchronize failed: {e}"))
.is_ok()
} else if !returned.registered {
true
} else {
*synced_devices.entry(returned.device).or_insert_with(|| {
set_device_by_id(returned.device)
.and_then(|_| device_synchronize())
.map_err(|e| {
tracing::debug!(
"synchronizing device {} failed: {e}",
returned.device
)
})
.is_ok()
})
};
if drained {
recycle(returned);
} else {
release(returned);
}
}
}
})
.expect("failed to spawn pinned-cleaner thread");
Cleaner { tx, thread }
}
pub fn shutdown() {
let cleaner = cleaner_slot().lock().unwrap().take();
if let Some(Cleaner { tx, thread }) = cleaner {
drop(tx);
let _ = thread.join();
}
clear();
}
fn recycle(mut returned: Returned) {
if !returned.registered {
if !register_region(returned.data.as_mut_ptr(), returned.data.len()) {
return; }
returned.registered = true;
}
let dirty_len = returned.dirty_len.min(returned.data.len());
returned.data[..dirty_len].fill(0);
let size = returned.data.len();
let mut evicted;
{
let mut pool = pool().lock().unwrap();
evicted = evict_to_fit(&mut pool, size);
if pool.total_bytes + size <= pool.max_bytes {
pool.total_bytes += size;
pool.by_size.entry(size).or_default().push(returned.data);
} else {
evicted.push(returned.data);
}
}
for mut buf in evicted {
unregister_region(buf.as_mut_ptr());
}
}
pub fn take(min_size: usize) -> PinnedBuffer {
let size = min_size.next_power_of_two();
{
let mut pool = pool().lock().unwrap();
if let Some(bufs) = pool.by_size.get_mut(&size) {
let data = bufs.pop().expect("empty size classes are removed");
if bufs.is_empty() {
pool.by_size.remove(&size);
}
pool.total_bytes -= size;
debug_assert_eq!(data.len(), size);
return PinnedBuffer {
data,
dirty_len: size,
registered: true,
last_use: Vec::new(),
};
}
}
PinnedBuffer {
data: vec![0u8; size],
dirty_len: size,
registered: false,
last_use: Vec::new(),
}
}
pub fn clear() {
let mut pool = pool().lock().unwrap();
for (_, bufs) in pool.by_size.iter_mut() {
for mut buf in bufs.drain(..) {
unregister_region(buf.as_mut_ptr());
}
}
pool.by_size.clear();
pool.total_bytes = 0;
}
#[cfg(test)]
mod tests {
use std::time::{Duration, Instant};
use super::*;
static TEST_LOCK: Mutex<()> = Mutex::new(());
fn lock_tests() -> std::sync::MutexGuard<'static, ()> {
TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner())
}
fn wait_for_pooled(size: usize) {
let deadline = Instant::now() + Duration::from_secs(30);
while pool()
.lock()
.unwrap()
.by_size
.get(&size)
.is_none_or(|bufs| bufs.is_empty())
{
assert!(
Instant::now() < deadline,
"size-{size} buffer never came back to the pool"
);
std::thread::sleep(Duration::from_millis(10));
}
}
#[test]
fn take_rounds_up_to_next_power_of_two_and_zero_fills() {
let _lock = lock_tests();
for (min_size, expected) in [(1, 1), (3, 4), (1024, 1024), (1025, 2048)] {
let buf = take(min_size);
assert_eq!(buf.len(), expected, "take({min_size})");
assert!(buf.iter().all(|&b| b == 0), "take({min_size}) not zeroed");
}
}
#[test]
fn round_trip_recycles_registered_rezeroed_buffer() {
let _lock = lock_tests();
const SIZE: usize = 1 << 13;
let mut buf = take(SIZE);
let ptr = buf.as_ptr() as usize;
buf.fill(0xAB);
buf.set_dirty_len(usize::MAX);
drop(buf);
wait_for_pooled(SIZE);
let buf = take(SIZE);
assert!(buf.is_pinned(), "pool hit should be page-locked");
assert_eq!(
buf.as_ptr() as usize,
ptr,
"pool hit should reuse the allocation"
);
assert!(buf.iter().all(|&b| b == 0), "recycled buffer not re-zeroed");
}
#[test]
fn shutdown_drains_joins_and_respawns_on_demand() {
let _lock = lock_tests();
const SIZE: usize = 1 << 15;
let mut buf = take(SIZE);
buf.fill(0xEF);
drop(buf); shutdown();
assert!(
pool().lock().unwrap().by_size.is_empty(),
"shutdown left buffers pooled"
);
let buf = take(SIZE);
assert!(!buf.is_pinned(), "pool should be empty after shutdown");
drop(buf);
wait_for_pooled(SIZE);
assert!(take(SIZE).is_pinned());
}
#[test]
fn byte_cap_evicts_largest_first() {
let _lock = lock_tests();
const SIZE: usize = 1 << 16;
const MARKER: usize = 1 << 10;
shutdown();
set_max_pooled_bytes(2 * SIZE + MARKER);
let bufs: Vec<_> = (0..3).map(|_| take(SIZE)).collect();
drop(bufs);
drop(take(MARKER));
wait_for_pooled(MARKER);
{
let pool = pool().lock().unwrap();
assert_eq!(
pool.by_size.get(&SIZE).map(|bufs| bufs.len()),
Some(2),
"third buffer should have evicted one of the first two"
);
assert_eq!(pool.total_bytes, 2 * SIZE + MARKER);
}
set_max_pooled_bytes(DEFAULT_MAX_POOLED_BYTES);
}
#[test]
fn recorded_last_use_event_gates_reuse() {
let _lock = lock_tests();
const SIZE: usize = 1 << 14;
let stream = crate::stream::CudaStream::new_non_blocking().unwrap();
let mut buf = take(SIZE);
let ptr = buf.as_ptr() as usize;
buf.fill(0xCD);
buf.record_last_use(&stream).unwrap();
drop(buf);
wait_for_pooled(SIZE);
let buf = take(SIZE);
assert!(buf.is_pinned(), "pool hit should be page-locked");
assert_eq!(
buf.as_ptr() as usize,
ptr,
"pool hit should reuse the allocation"
);
assert!(buf.iter().all(|&b| b == 0), "recycled buffer not re-zeroed");
}
}