use std::cell::{Cell, RefCell};
use std::collections::BTreeMap;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
pub struct SizeClass(usize);
impl SizeClass {
#[must_use]
pub fn for_size(size: usize) -> Self {
if size == 0 {
return Self(0);
}
let bits = usize::BITS - (size - 1).leading_zeros();
let result = Self(1 << bits);
debug_assert!(
result.0.is_power_of_two(),
"size class must be a power of 2"
);
debug_assert!(result.0 >= size, "size class must be >= requested size");
result
}
#[must_use]
pub const fn allocation_size(&self) -> usize {
self.0
}
}
#[derive(Debug)]
pub struct PooledBuffer {
data: Vec<f32>,
size_class: SizeClass,
len: usize,
}
impl PooledBuffer {
fn new(size_class: SizeClass) -> Self {
Self {
data: vec![0.0; size_class.allocation_size()],
size_class,
len: 0,
}
}
#[must_use]
pub fn as_slice(&self) -> &[f32] {
&self.data[..self.len]
}
pub fn as_mut_slice(&mut self) -> &mut [f32] {
&mut self.data[..self.len]
}
#[must_use]
pub fn capacity(&self) -> usize {
self.data.len()
}
#[must_use]
pub fn len(&self) -> usize {
self.len
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len == 0
}
pub fn set_len(&mut self, len: usize) {
assert!(len <= self.capacity(), "length exceeds capacity");
self.len = len;
}
pub fn fill(&mut self, value: f32) {
for v in &mut self.data[..self.len] {
*v = value;
}
}
#[must_use]
pub const fn size_class(&self) -> SizeClass {
self.size_class
}
#[must_use]
pub fn into_vec(mut self) -> Vec<f32> {
self.data.truncate(self.len);
self.data
}
pub fn copy_from_slice(&mut self, src: &[f32]) {
assert!(src.len() <= self.capacity(), "source too large");
self.data[..src.len()].copy_from_slice(src);
self.len = src.len();
}
}
#[derive(Debug, Default)]
struct PoolCounters {
allocations: Cell<usize>,
hits: Cell<usize>,
misses: Cell<usize>,
returns: Cell<usize>,
dropped: Cell<usize>,
}
impl PoolCounters {
fn inc(counter: &Cell<usize>) {
counter.set(counter.get() + 1);
}
fn to_stats(&self) -> PoolStats {
PoolStats {
allocations: self.allocations.get(),
hits: self.hits.get(),
misses: self.misses.get(),
returns: self.returns.get(),
dropped: self.dropped.get(),
}
}
}
#[derive(Debug)]
pub struct MemoryPool {
pools: RefCell<BTreeMap<SizeClass, Vec<PooledBuffer>>>,
counters: PoolCounters,
max_per_class: usize,
}
#[derive(Debug, Default, Clone)]
pub struct PoolStats {
pub allocations: usize,
pub hits: usize,
pub misses: usize,
pub returns: usize,
pub dropped: usize,
}
impl PoolStats {
#[must_use]
pub fn hit_rate(&self) -> f32 {
if self.allocations == 0 {
0.0
} else {
(self.hits as f32) / (self.allocations as f32) * 100.0
}
}
}
impl Default for MemoryPool {
fn default() -> Self {
Self::new()
}
}
impl MemoryPool {
#[must_use]
pub fn new() -> Self {
Self::with_max_per_class(16)
}
#[must_use]
pub fn with_max_per_class(max_per_class: usize) -> Self {
Self {
pools: RefCell::new(BTreeMap::new()),
counters: PoolCounters::default(),
max_per_class,
}
}
pub fn get(&self, size: usize) -> PooledBuffer {
let size_class = SizeClass::for_size(size);
let cached = self
.pools
.borrow_mut()
.get_mut(&size_class)
.and_then(Vec::pop);
PoolCounters::inc(&self.counters.allocations);
if let Some(mut buffer) = cached {
PoolCounters::inc(&self.counters.hits);
buffer.set_len(size);
buffer.fill(0.0);
buffer
} else {
PoolCounters::inc(&self.counters.misses);
let mut buffer = PooledBuffer::new(size_class);
buffer.set_len(size);
buffer
}
}
pub fn get_from_slice(&self, data: &[f32]) -> PooledBuffer {
let mut buffer = self.get(data.len());
buffer.copy_from_slice(data);
buffer
}
pub fn return_buffer(&self, buffer: PooledBuffer) {
PoolCounters::inc(&self.counters.returns);
let mut pools = self.pools.borrow_mut();
let pool = pools.entry(buffer.size_class).or_default();
if pool.len() < self.max_per_class {
pool.push(buffer);
} else {
PoolCounters::inc(&self.counters.dropped);
}
}
#[must_use]
pub fn stats(&self) -> PoolStats {
self.counters.to_stats()
}
pub fn clear(&self) {
self.pools.borrow_mut().clear();
}
#[must_use]
pub fn buffered_count(&self) -> usize {
self.with_pools(|pools| pools.values().map(Vec::len).sum())
}
#[must_use]
pub fn buffered_bytes(&self) -> usize {
self.with_pools(|pools| {
pools
.iter()
.map(|(class, buffers)| class.allocation_size() * buffers.len() * 4)
.sum()
})
}
fn with_pools<F, R>(&self, f: F) -> R
where
F: FnOnce(&BTreeMap<SizeClass, Vec<PooledBuffer>>) -> R,
{
f(&self.pools.borrow())
}
pub fn preallocate(&self, sizes: &[usize]) {
for &size in sizes {
let buffer = self.get(size);
self.return_buffer(buffer);
}
}
}
thread_local! {
static POOL: MemoryPool = MemoryPool::new();
}
#[must_use]
pub fn get_buffer(size: usize) -> PooledBuffer {
POOL.with(|pool| pool.get(size))
}
pub fn return_buffer(buffer: PooledBuffer) {
POOL.with(|pool| pool.return_buffer(buffer));
}
#[must_use]
pub fn pool_stats() -> PoolStats {
POOL.with(|pool| pool.stats())
}
#[cfg(test)]
#[path = "pool_tests.rs"]
mod tests;