use std::collections::HashMap;
use std::ffi::c_void;
use std::sync::Mutex;
use super::context::Api;
pub(super) const CLASS_CAP: usize = 256 * 1024 * 1024;
const CLASS_FLOOR: usize = 4096;
pub(super) const PARKED_CAP: usize = 4 * CLASS_CAP;
pub(super) struct Pool {
parked: Mutex<Parked>,
}
struct Parked {
classes: HashMap<usize, Vec<*mut c_void>>,
bytes: usize,
}
impl Pool {
pub(super) fn new() -> Self {
Self {
parked: Mutex::new(Parked {
classes: HashMap::new(),
bytes: 0,
}),
}
}
pub(super) fn take(&self, api: &Api, bytes: usize) -> Result<*mut c_void, String> {
let class = class_of(bytes);
{
let mut parked = self.parked.lock().expect("the pool mutex is poisoned");
if let Some(buffer) = parked.classes.get_mut(&class).and_then(Vec::pop) {
parked.bytes -= class;
return Ok(buffer);
}
}
let mut buffer = std::ptr::null_mut();
let status = unsafe { (api.malloc)(&mut buffer, class) };
if status != 0 {
return Err(format!("cudaMalloc failed: {}", api.error_string(status)));
}
Ok(buffer)
}
pub(super) fn give(&self, api: &Api, bytes: usize, buffer: *mut c_void) {
let class = class_of(bytes);
if class <= CLASS_CAP {
let mut parked = self.parked.lock().expect("the pool mutex is poisoned");
if parked.bytes + class <= PARKED_CAP {
parked.classes.entry(class).or_default().push(buffer);
parked.bytes += class;
return;
}
}
let _ = unsafe { (api.free)(buffer) };
}
}
fn class_of(bytes: usize) -> usize {
if bytes > CLASS_CAP {
return bytes;
}
bytes.next_power_of_two().max(CLASS_FLOOR)
}