use super::{ManagedMemoryHandle, MemoryPool, PageMapping, Slice, calculate_padding};
use crate::memory_management::{BytesFormat, MemoryLocation, MemoryPoolKind, MemoryPoolReport};
use crate::storage::StorageUtilization;
use crate::{memory_management::MemoryUsage, server::IoError};
use alloc::vec::Vec;
use cubecl_environment::backtrace::BackTrace;
pub struct DirectPool {
slices: Vec<Option<Slice>>,
vacant: Vec<usize>,
alignment: u64,
location_base: MemoryLocation,
reclaim_at: Option<u64>,
pages_peak: u64,
largest_alloc: u64,
}
impl DirectPool {
pub fn new(alignment: u64, pool_pos: u8, reclaim_at: Option<u64>) -> Self {
Self {
slices: Vec::new(),
vacant: Vec::new(),
alignment,
location_base: MemoryLocation::new(pool_pos, 0, 0),
reclaim_at,
pages_peak: 0,
largest_alloc: 0,
}
}
pub(crate) fn report(&self) -> MemoryPoolReport {
MemoryPoolReport {
kind: MemoryPoolKind::Direct,
usage: self.get_memory_usage(),
pages: self.live().count() as u64,
pages_peak: self.pages_peak,
pages_unmapped: self.live().filter(|slice| !slice.mapped).count() as u64,
largest_alloc: self.largest_alloc,
}
}
fn live(&self) -> impl Iterator<Item = &Slice> {
self.slices.iter().flatten()
}
fn reserved(&self) -> u64 {
self.live().map(|slice| slice.effective_size()).sum()
}
fn release_free<Storage: crate::storage::ComputeStorage>(
&mut self,
storage: &mut Storage,
headroom: u64,
) {
let Some(ceiling) = self.reclaim_at else {
return;
};
let mut reserved = self.reserved();
if reserved + headroom <= ceiling {
return;
}
for (index, entry) in self.slices.iter_mut().enumerate() {
if reserved + headroom <= ceiling {
break;
}
let Some(slice) = entry else { continue };
if !slice.is_free() {
continue;
}
if slice.mapped {
storage.dealloc(slice.storage.id);
}
reserved -= slice.effective_size();
*entry = None;
self.vacant.push(index);
}
}
}
impl MemoryPool for DirectPool {
fn accept(&self, _size: u64) -> bool {
true
}
fn find(&self, binding: &super::ManagedMemoryBinding) -> Result<&Slice, IoError> {
let index = binding.descriptor().slice();
self.slices
.get(index)
.and_then(|slice| slice.as_ref())
.ok_or_else(|| IoError::NotFound {
backtrace: BackTrace::capture(),
reason: alloc::format!("Memory slice {index} doesn't exist").into(),
})
}
fn try_reserve(&mut self, size: u64) -> Option<ManagedMemoryHandle> {
let padding = calculate_padding(size, self.alignment);
let effective_size = size + padding;
let slice = self
.slices
.iter_mut()
.flatten()
.find(|slice| slice.is_free() && slice.effective_size() == effective_size)?;
slice.padding = padding;
slice.storage.utilization = StorageUtilization { offset: 0, size };
self.largest_alloc = self.largest_alloc.max(size);
Some(slice.handle.clone())
}
fn alloc<Storage: crate::storage::ComputeStorage>(
&mut self,
storage: &mut Storage,
size: u64,
mapping: PageMapping,
) -> Result<ManagedMemoryHandle, IoError> {
let padding = calculate_padding(size, self.alignment);
let effective_size = size + padding;
self.release_free(storage, effective_size);
let storage_handle = mapping.storage_handle(storage, effective_size)?;
let mut slice = Slice::new(storage_handle, padding);
slice.mapped = matches!(mapping, PageMapping::Eager);
slice.storage.utilization = StorageUtilization { offset: 0, size };
let index = match self.vacant.pop() {
Some(index) => index,
None => {
self.slices.push(None);
self.slices.len() - 1
}
};
let mut location = self.location_base;
location.slice = index as u32;
slice.descriptor().update_location(location);
let handle = slice.handle.clone();
self.slices[index] = Some(slice);
self.pages_peak = self.pages_peak.max(self.live().count() as u64);
self.largest_alloc = self.largest_alloc.max(size);
Ok(handle)
}
fn materialize<Storage: crate::storage::ComputeStorage>(
&mut self,
storage: &mut Storage,
binding: &super::ManagedMemoryBinding,
) -> Result<(), IoError> {
let index = binding.descriptor().slice();
let Some(slice) = self.slices.get_mut(index).and_then(|slice| slice.as_mut()) else {
return Ok(());
};
if slice.mapped || slice.handle.descriptor() != binding.descriptor() {
return Ok(());
}
slice.materialize(storage)
}
fn get_memory_usage(&self) -> MemoryUsage {
let used: Vec<_> = self.live().filter(|slice| !slice.is_free()).collect();
MemoryUsage {
number_allocs: used.len() as u64,
bytes_in_use: used.iter().map(|slice| slice.storage.size()).sum(),
bytes_padding: used.iter().map(|slice| slice.padding).sum(),
bytes_reserved: self.live().map(|slice| slice.effective_size()).sum(),
}
}
fn cleanup<Storage: crate::storage::ComputeStorage>(
&mut self,
storage: &mut Storage,
_alloc_nr: u64,
explicit: bool,
) {
if !explicit {
return;
}
for (index, entry) in self.slices.iter_mut().enumerate() {
let Some(slice) = entry else { continue };
if !slice.is_free() {
continue;
}
if slice.mapped {
storage.dealloc(slice.storage.id);
}
*entry = None;
self.vacant.push(index);
}
}
fn bind(
&mut self,
reserved: ManagedMemoryHandle,
assigned: ManagedMemoryHandle,
_cursor: u64,
) -> Result<(), IoError> {
let index = reserved.descriptor().slice();
let slice = self
.slices
.get_mut(index)
.and_then(|slice| slice.as_mut())
.ok_or_else(|| IoError::NotFound {
backtrace: BackTrace::capture(),
reason: alloc::format!("Memory slice {index} doesn't exist").into(),
})?;
assigned
.descriptor()
.update_location(reserved.descriptor().location());
slice.handle = assigned;
Ok(())
}
}
impl core::fmt::Display for DirectPool {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
let usage = self.get_memory_usage();
if usage.bytes_reserved == 0 {
return Ok(());
}
f.write_fmt(format_args!(
" - Direct: {} slices, largest {}\n",
self.live().count(),
BytesFormat::new(self.largest_alloc)
))?;
f.write_fmt(format_args!("\n{usage}\n"))
}
}