use std::{
cmp::max, collections::BTreeMap, error::Error, fmt, marker::PhantomData, mem::size_of,
sync::{Arc, Mutex},
};
use cudarc::driver::{
sys::CUdeviceptr, CudaSlice, CudaStream, DevicePtr, DevicePtrMut, DeviceSlice, DriverError, SyncOnDrop,
};
pub type SegId = usize;
const MIN_SEG: usize = 1 << 20;
#[derive(Debug)]
pub enum AllocError {
Cuda(DriverError),
InvalidRequest { bytes: usize },
}
impl fmt::Display for AllocError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
AllocError::Cuda(e) => write!(f, "cuda alloc failed: {e}"),
AllocError::InvalidRequest { bytes } => write!(f, "invalid alloc request of {bytes} bytes"),
}
}
}
impl Error for AllocError {}
impl From<DriverError> for AllocError {
fn from(e: DriverError) -> Self {
AllocError::Cuda(e)
}
}
#[derive(Clone)]
pub struct CachingAllocator {
pool: Arc<Mutex<Pool>>,
stream: Arc<CudaStream>,
}
pub struct CudaBuffer<T> {
storage: Arc<CudaSlice<u8>>, offset: usize, len: usize, bytes: usize, seg: SegId, pool: Arc<Mutex<Pool>>, _t: PhantomData<T>,
}
pub struct Pool {
stream: Arc<CudaStream>,
segments: Vec<Option<Segment>>,
allocated: usize,
reserved: usize,
active: usize,
pool_miss: usize,
peak: usize,
}
#[derive(Debug, Clone, Copy, Default)]
pub struct Stats {
pub allocated: usize, pub reserved: usize, pub active: usize, pub pool_miss: usize, pub peak: usize, }
pub struct Segment {
storage: Arc<CudaSlice<u8>>,
total: usize,
live: usize,
free: BTreeMap<usize, usize>,
}
impl CachingAllocator {
pub fn new(stream: &Arc<CudaStream>) -> Self {
CachingAllocator {
pool: Arc::new(Mutex::new(Pool::new(stream))),
stream: stream.clone(),
}
}
pub fn stream(&self) -> &Arc<CudaStream> {
&self.stream
}
pub fn alloc<T>(&self, len: usize) -> Result<CudaBuffer<T>, AllocError> {
let bytes = len
.checked_mul(size_of::<T>())
.ok_or(AllocError::InvalidRequest { bytes: 0 })?;
if bytes == 0 {
return Err(AllocError::InvalidRequest { bytes: 0 });
}
let (storage, seg, offset) = self.pool.lock().expect("pool poisoned").alloc_raw(bytes)?;
Ok(CudaBuffer {
storage,
offset,
len,
bytes: align_up(bytes, 256),
seg,
pool: self.pool.clone(),
_t: PhantomData,
})
}
pub fn empty_cache(&self) {
self.pool.lock().expect("pool poisoned").empty_cache();
}
pub fn stats(&self) -> Stats {
self.pool.lock().expect("pool poisoned").stats()
}
pub fn assert_consistency(&self) {
self.pool.lock().expect("pool poisoned").assert_consistency();
}
}
impl<T> CudaBuffer<T> {
pub fn offset(&self) -> usize {
self.offset
}
pub fn byte_len(&self) -> usize {
self.bytes
}
}
impl<T> Drop for CudaBuffer<T> {
fn drop(&mut self) {
let _ = self.storage.stream().synchronize();
if let Ok(mut g) = self.pool.lock() {
g.free_raw(self.seg, self.offset, self.bytes);
}
}
}
impl<T> DeviceSlice<T> for CudaBuffer<T> {
fn len(&self) -> usize {
self.len
}
fn stream(&self) -> &Arc<CudaStream> {
self.storage.stream()
}
}
impl<T> DevicePtr<T> for CudaBuffer<T> {
fn device_ptr<'a>(&'a self, stream: &'a CudaStream) -> (CUdeviceptr, SyncOnDrop<'a>) {
let (base, guard) = <CudaSlice<u8> as DevicePtr<u8>>::device_ptr(&self.storage, stream);
(base + self.offset as u64, guard)
}
}
impl<T> DevicePtrMut<T> for CudaBuffer<T> {
fn device_ptr_mut<'a>(&'a mut self, stream: &'a CudaStream) -> (CUdeviceptr, SyncOnDrop<'a>) {
DevicePtr::device_ptr(self, stream)
}
}
impl Pool {
pub fn new(stream: &Arc<CudaStream>) -> Self {
Pool {
stream: stream.clone(),
segments: Vec::new(),
allocated: 0,
reserved: 0,
active: 0,
pool_miss: 0,
peak: 0,
}
}
pub fn stats(&self) -> Stats {
Stats {
allocated: self.allocated,
reserved: self.reserved,
active: self.active,
pool_miss: self.pool_miss,
peak: self.peak,
}
}
pub fn alloc_raw(&mut self, bytes: usize) -> Result<(Arc<CudaSlice<u8>>, SegId, usize), AllocError> {
let bytes = align_up(bytes, 256);
if bytes == 0 {
return Err(AllocError::InvalidRequest { bytes: 0 });
}
let mut best_sel: Option<(SegId, usize, usize)> = None;
for (i, slot) in self.segments.iter().enumerate() {
if let Some(seg) = slot {
if let Some((off, run_len)) = seg.find_best(bytes) {
best_sel = match best_sel {
Some((_, _, last_len)) if last_len < run_len => best_sel,
_ => Some((i, off, run_len)),
};
}
}
}
let (seg_id, offset) = match best_sel {
Some((id, off, _)) => (id, off),
None => (self.new_segment(bytes)?, 0),
};
let offset = self.segments[seg_id]
.as_mut()
.unwrap()
.alloc_raw_in(bytes, offset)?;
let slice = self.segments[seg_id].as_ref().unwrap().storage().clone();
self.allocated += bytes;
self.peak = self.peak.max(self.allocated);
Ok((slice, seg_id, offset))
}
fn new_segment(&mut self, bytes: usize) -> Result<SegId, AllocError> {
let seg_bytes = align_up(max(bytes, MIN_SEG), MIN_SEG);
let seg = Segment::new(seg_bytes, &self.stream)?;
self.pool_miss += 1;
self.reserved += seg.total;
self.active += 1;
self.segments.push(Some(seg));
Ok(self.segments.len() - 1)
}
pub fn free_raw(&mut self, id: SegId, offset: usize, bytes: usize) {
assert!(self.segments[id].is_some(), "free on dead segment {id}");
let bytes = align_up(bytes, 256);
assert!(self.allocated >= bytes, "allocated underflow on free");
self.segments[id].as_mut().unwrap().free_raw(offset, bytes);
self.allocated -= bytes;
}
pub fn empty_cache(&mut self) {
for slot in self.segments.iter_mut() {
if let Some(seg) = slot {
if seg.live == 0 {
*slot = None;
}
}
}
self.recount();
}
fn recount(&mut self) {
self.allocated = self.segments.iter().flatten().map(|s| s.live).sum();
self.reserved = self.segments.iter().flatten().map(|s| s.total).sum();
self.active = self.segments.iter().filter(|s| s.is_some()).count();
}
pub fn assert_consistency(&self) {
for seg in self.segments.iter().flatten() {
seg.assert_consistency();
}
let sum_live: usize = self.segments.iter().flatten().map(|s| s.live).sum();
let sum_total: usize = self.segments.iter().flatten().map(|s| s.total).sum();
let n_active = self.segments.iter().filter(|s| s.is_some()).count();
assert_eq!(self.allocated, sum_live, "pool allocated != Σlive");
assert_eq!(self.reserved, sum_total, "pool reserved != Σtotal");
assert_eq!(self.active, n_active, "pool active != count(Some)");
}
}
impl Segment {
pub fn new(total: usize, stream: &Arc<CudaStream>) -> Result<Self, AllocError> {
let total = align_up(total.max(1), 256);
let storage = Arc::new(unsafe { stream.alloc::<u8>(total)? });
let mut free = BTreeMap::new();
free.insert(0, total);
Ok(Segment { storage, total, live: 0, free })
}
pub fn total(&self) -> usize {
self.total
}
pub fn live(&self) -> usize {
self.live
}
pub fn storage(&self) -> &Arc<CudaSlice<u8>> {
&self.storage
}
pub fn find_best(&self, bytes: usize) -> Option<(usize, usize)> {
assert!(bytes % 256 == 0);
let mut best: Option<(usize, usize)> = None;
for (&offset, &size) in self.free.iter() {
if size >= bytes {
best = match best {
Some((_, last)) if last < size => best, _ => Some((offset, size)),
};
}
}
best
}
pub fn alloc_raw_in(&mut self, bytes: usize, offset: usize) -> Result<usize, AllocError> {
assert!(bytes % 256 == 0);
if bytes == 0 || bytes > self.total {
return Err(AllocError::InvalidRequest { bytes });
}
let origin = self
.free
.get(&offset)
.copied()
.filter(|&size| size >= bytes)
.ok_or(AllocError::InvalidRequest { bytes })?;
self.free.remove(&offset);
let leftover = origin - bytes;
if leftover > 0 {
self.free.insert(offset + bytes, leftover);
}
self.live += bytes;
Ok(offset)
}
pub fn free_raw(&mut self, offset: usize, bytes: usize) {
assert!(offset % 256 == 0, "offset must be 256-aligned");
let bytes = align_up(bytes, 256);
assert!(bytes > 0, "cannot free 0 bytes");
assert!(offset + bytes <= self.total, "free range out of segment");
assert!(self.live >= bytes, "live underflow / double free");
let mut start = offset;
if let Some((&k, &v)) = self.free.range(..offset).next_back() {
assert!(k + v <= offset, "double free: overlaps left free run");
if k + v == offset {
start = k;
}
}
let mut end = offset + bytes;
if let Some((&k, &v)) = self.free.range(offset..).next() {
assert!(k >= offset + bytes, "double free: overlaps right free run");
if k == offset + bytes {
end = k + v;
}
}
self.free.retain(|&k, _| k < start || k >= end);
self.free.insert(start, end - start);
self.live -= bytes;
}
pub fn assert_consistency(&self) {
let mut sum_free = 0usize;
let mut prev_key: Option<usize> = None;
for (&k, &v) in self.free.iter() {
assert!(v > 0, "zero-length free run at {k}");
assert!(k + v <= self.total, "free run [{k}, {}) exceeds total {}", k + v, self.total);
if let Some(pk) = prev_key {
assert!(
pk + self.free[&pk] < k,
"free runs at {pk} and {k} are touching (coalesce missed)"
);
}
sum_free += v;
prev_key = Some(k);
}
assert_eq!(self.live + sum_free, self.total, "conservation broken: live+free != total");
}
}
fn align_up(v: usize, a: usize) -> usize {
(v + a - 1) & !(a - 1)
}
#[cfg(test)]
mod tests {
use super::*;
use cudarc::driver::CudaContext;
fn allocator() -> CachingAllocator {
CachingAllocator::new(&CudaContext::new(0).unwrap().default_stream())
}
#[test]
fn two_buffers_same_segment() {
let alloc = allocator();
let stream = alloc.stream().clone();
let n = 64 * 1024;
let mut a = alloc.alloc::<f32>(n).unwrap();
let mut b = alloc.alloc::<f32>(n).unwrap();
assert_eq!(a.offset(), 0);
assert_eq!(b.offset(), align_up(n * 4, 256));
assert_ne!(a.offset(), b.offset());
let ha: Vec<f32> = (0..n).map(|i| i as f32).collect();
let hb: Vec<f32> = (0..n).map(|i| -(i as f32)).collect();
stream.memcpy_htod(&ha, &mut a).unwrap();
stream.memcpy_htod(&hb, &mut b).unwrap();
stream.synchronize().unwrap();
assert_eq!(stream.clone_dtoh(&a).unwrap(), ha);
assert_eq!(stream.clone_dtoh(&b).unwrap(), hb);
alloc.assert_consistency();
drop(a);
drop(b);
alloc.empty_cache();
assert_eq!((alloc.stats().reserved, alloc.stats().active), (0, 0));
}
#[test]
fn buffer_reuse_after_drop() {
let alloc = allocator();
let n = 256 * 1024;
let off0 = alloc.alloc::<f32>(n).unwrap().offset();
drop(alloc.alloc::<f32>(n).unwrap());
let off1 = alloc.alloc::<f32>(n).unwrap().offset();
assert_eq!(off0, off1, "freed block must be reused");
assert_eq!(alloc.stats().pool_miss, 1, "single segment serves everything");
alloc.assert_consistency();
}
#[test]
fn cross_type_reuse() {
let alloc = allocator();
let stream = alloc.stream().clone();
let nf = 64 * 1024; let nb = nf * 4;
let mut f1 = alloc.alloc::<f32>(nf).unwrap();
let off_f = f1.offset();
let host1 = vec![1.5f32; nf];
stream.memcpy_htod(&host1, &mut f1).unwrap();
stream.synchronize().unwrap();
assert_eq!(stream.clone_dtoh(&f1).unwrap(), host1);
drop(f1);
let mut b = alloc.alloc::<u8>(nb).unwrap();
assert_eq!(b.offset(), off_f, "u8 must reuse the freed f32 bytes");
let hostb: Vec<u8> = (0..nb).map(|i| (i % 251) as u8).collect();
stream.memcpy_htod(&hostb, &mut b).unwrap();
stream.synchronize().unwrap();
assert_eq!(stream.clone_dtoh(&b).unwrap(), hostb);
drop(b);
let mut f2 = alloc.alloc::<f32>(nf).unwrap();
assert_eq!(f2.offset(), off_f, "f32 must reuse the same bytes again");
let host2 = vec![42.0f32; nf];
stream.memcpy_htod(&host2, &mut f2).unwrap();
stream.synchronize().unwrap();
assert_eq!(stream.clone_dtoh(&f2).unwrap(), host2, "no stale data from u8 phase");
drop(f2);
alloc.assert_consistency();
}
#[test]
fn zero_request_is_error() {
let alloc = allocator();
assert!(matches!(alloc.alloc::<f32>(0), Err(AllocError::InvalidRequest { .. })));
}
}