use std::any::{Any, TypeId};
use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Mutex;
use oxicuda_memory::{MemoryPool, PooledBuffer};
use crate::NeuralResult;
use sklears_core::error::SklearsError;
#[derive(Debug)]
pub struct PoolTelemetry {
hits: AtomicU64,
misses: AtomicU64,
allocated_bytes: AtomicU64,
capacity_bytes: u64,
}
impl PoolTelemetry {
pub fn new(capacity_bytes: usize) -> Self {
Self {
hits: AtomicU64::new(0),
misses: AtomicU64::new(0),
allocated_bytes: AtomicU64::new(0),
capacity_bytes: capacity_bytes as u64,
}
}
pub fn record_hit(&self) {
self.hits.fetch_add(1, Ordering::Relaxed);
}
pub fn record_miss(&self, bytes: usize) {
self.misses.fetch_add(1, Ordering::Relaxed);
self.allocated_bytes
.fetch_add(bytes as u64, Ordering::Relaxed);
}
pub fn record_free(&self, bytes: usize) {
let _ =
self.allocated_bytes
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| {
Some(current.saturating_sub(bytes as u64))
});
}
pub fn hit_rate(&self) -> f64 {
let hits = self.hits.load(Ordering::Relaxed) as f64;
let misses = self.misses.load(Ordering::Relaxed) as f64;
let total = hits + misses;
if total == 0.0 {
0.0
} else {
hits / total
}
}
pub fn used_fraction(&self) -> f64 {
if self.capacity_bytes == 0 {
return 0.0;
}
let used = self.allocated_bytes.load(Ordering::Relaxed) as f64;
(used / self.capacity_bytes as f64).clamp(0.0, 1.0)
}
}
type PoolFreeLists = Mutex<HashMap<(TypeId, usize), Vec<Box<dyn Any>>>>;
pub struct GpuMemoryPool {
inner_pool: MemoryPool,
free_lists: PoolFreeLists,
telemetry: PoolTelemetry,
}
impl GpuMemoryPool {
const MAX_CACHED_PER_CLASS: usize = 4;
pub fn new(device_ordinal: i32, capacity_bytes: usize) -> NeuralResult<Self> {
let inner_pool = MemoryPool::new(device_ordinal).map_err(|e| {
SklearsError::InvalidInput(format!("Failed to create GPU memory pool: {}", e))
})?;
Ok(Self {
inner_pool,
free_lists: Mutex::new(HashMap::new()),
telemetry: PoolTelemetry::new(capacity_bytes),
})
}
pub fn acquire<T: Copy + 'static>(
&self,
n: usize,
stream: &oxicuda_driver::Stream,
) -> NeuralResult<PooledHandle<'_, T>> {
if n == 0 {
return Err(SklearsError::InvalidInput(
"cannot acquire a zero-element pooled GPU buffer".to_string(),
));
}
let size_class = n.checked_next_power_of_two().unwrap_or(n);
let key = (TypeId::of::<T>(), size_class);
let reused = {
let mut lists = self
.free_lists
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
lists.get_mut(&key).and_then(|bucket| bucket.pop())
};
if let Some(boxed) = reused {
match boxed.downcast::<PooledBuffer<T>>() {
Ok(buf) => {
self.telemetry.record_hit();
return Ok(PooledHandle {
buf: Some(*buf),
pool: self,
size_class,
});
}
Err(_mismatched) => {
}
}
}
let bytes = size_class
.checked_mul(std::mem::size_of::<T>())
.ok_or_else(|| {
SklearsError::InvalidInput("pooled allocation size overflow".to_string())
})?;
let buf =
PooledBuffer::<T>::alloc_async(&self.inner_pool, size_class, stream).map_err(|e| {
SklearsError::InvalidInput(format!("GPU pooled allocation failed: {}", e))
})?;
self.telemetry.record_miss(bytes);
Ok(PooledHandle {
buf: Some(buf),
pool: self,
size_class,
})
}
fn release<T: Copy + 'static>(&self, buf: PooledBuffer<T>, size_class: usize) {
let key = (TypeId::of::<T>(), size_class);
let mut lists = self
.free_lists
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let bucket = lists.entry(key).or_default();
if bucket.len() < Self::MAX_CACHED_PER_CLASS {
bucket.push(Box::new(buf));
} else {
drop(lists); let bytes = size_class.saturating_mul(std::mem::size_of::<T>());
self.telemetry.record_free(bytes);
drop(buf);
}
}
pub fn telemetry_stats(&self) -> (f64, f64) {
(self.telemetry.used_fraction(), self.telemetry.hit_rate())
}
}
pub struct PooledHandle<'a, T: Copy + 'static> {
buf: Option<PooledBuffer<T>>,
pool: &'a GpuMemoryPool,
size_class: usize,
}
impl<T: Copy + 'static> PooledHandle<'_, T> {
pub fn len(&self) -> usize {
self.buf.as_ref().map(|b| b.len()).unwrap_or(0)
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn byte_size(&self) -> usize {
self.buf.as_ref().map(|b| b.byte_size()).unwrap_or(0)
}
pub fn as_device_ptr(&self) -> oxicuda_driver::CUdeviceptr {
self.buf.as_ref().map(|b| b.as_device_ptr()).unwrap_or(0)
}
}
impl<T: Copy + 'static> Drop for PooledHandle<'_, T> {
fn drop(&mut self) {
if let Some(buf) = self.buf.take() {
self.pool.release(buf, self.size_class);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn hit_rate_matches_recorded_counts() {
let telemetry = PoolTelemetry::new(1024);
for _ in 0..3 {
telemetry.record_hit();
}
telemetry.record_miss(64);
assert!((telemetry.hit_rate() - 0.75).abs() < 1e-12);
}
#[test]
fn hit_rate_is_zero_not_nan_when_untouched() {
let telemetry = PoolTelemetry::new(1024);
assert_eq!(telemetry.hit_rate(), 0.0);
assert!(!telemetry.hit_rate().is_nan());
}
#[test]
fn used_fraction_is_zero_when_nothing_allocated() {
let telemetry = PoolTelemetry::new(1024);
assert_eq!(telemetry.used_fraction(), 0.0);
}
#[test]
fn used_fraction_reflects_partial_allocation() {
let telemetry = PoolTelemetry::new(1000);
telemetry.record_miss(250);
assert!((telemetry.used_fraction() - 0.25).abs() < 1e-12);
}
#[test]
fn used_fraction_is_clamped_when_allocation_exceeds_capacity() {
let telemetry = PoolTelemetry::new(100);
telemetry.record_miss(150);
assert_eq!(telemetry.used_fraction(), 1.0);
}
#[test]
fn used_fraction_with_zero_capacity_never_divides_by_zero() {
let telemetry = PoolTelemetry::new(0);
telemetry.record_miss(10);
assert_eq!(telemetry.used_fraction(), 0.0);
assert!(!telemetry.used_fraction().is_nan());
}
#[test]
fn record_free_reduces_allocated_bytes() {
let telemetry = PoolTelemetry::new(1000);
telemetry.record_miss(400);
assert!((telemetry.used_fraction() - 0.4).abs() < 1e-12);
telemetry.record_free(400);
assert_eq!(telemetry.used_fraction(), 0.0);
}
#[test]
fn record_free_saturates_instead_of_underflowing() {
let telemetry = PoolTelemetry::new(1000);
telemetry.record_miss(100);
telemetry.record_free(500); assert_eq!(telemetry.used_fraction(), 0.0);
}
use std::sync::Arc;
fn real_context() -> Option<Arc<oxicuda_driver::Context>> {
if oxicuda_driver::init().is_err() {
return None;
}
let device = oxicuda_driver::Device::get(0).ok()?;
oxicuda_driver::Context::new(&device).ok().map(Arc::new)
}
#[test]
fn acquire_after_release_is_a_genuine_hit_not_a_fresh_allocation() {
let Some(ctx) = real_context() else {
return;
};
let Ok(stream) = oxicuda_driver::Stream::new(&ctx) else {
return;
};
let Ok(pool) = GpuMemoryPool::new(0, 1_048_576) else {
return;
};
{
let handle = match pool.acquire::<f32>(16, &stream) {
Ok(handle) => handle,
Err(_) => return, };
assert_eq!(handle.len(), 16);
}
let (_, hit_rate_after_miss) = pool.telemetry_stats();
assert_eq!(hit_rate_after_miss, 0.0);
let Ok(_second) = pool.acquire::<f32>(16, &stream) else {
return;
};
let (_, hit_rate_after_hit) = pool.telemetry_stats();
assert_eq!(hit_rate_after_hit, 0.5); }
}