use std::fmt;
use std::sync::Arc;
use memmap2::Mmap;
use metal::Buffer as MetalBuffer;
use crate::dtypes::DType;
use crate::error::{MlxError, Result};
use crate::residency::ResidencySet;
pub struct MlxBuffer {
storage: Arc<MlxBufferStorage>,
dtype: DType,
shape: Vec<usize>,
byte_offset: u64,
data_byte_len: usize,
}
pub(crate) struct MlxBufferStorage {
inner: MetalBuffer,
residency_set: Option<ResidencySet>,
file_backing: Option<Arc<Mmap>>,
cpu_writable: bool,
}
impl Drop for MlxBufferStorage {
fn drop(&mut self) {
if let Some(set) = self.residency_set.as_ref() {
set.remove_allocation(&self.inner);
}
}
}
crate::static_assertions_send_sync!(MlxBuffer);
impl Clone for MlxBuffer {
fn clone(&self) -> Self {
Self {
storage: self.storage.clone(),
dtype: self.dtype,
shape: self.shape.clone(),
byte_offset: self.byte_offset,
data_byte_len: self.data_byte_len,
}
}
}
impl MlxBuffer {
pub fn from_raw(inner: MetalBuffer, dtype: DType, shape: Vec<usize>) -> Self {
let data_byte_len = inner.length() as usize;
Self {
storage: Arc::new(MlxBufferStorage {
inner,
residency_set: None,
file_backing: None,
cpu_writable: true,
}),
dtype,
shape,
byte_offset: 0,
data_byte_len,
}
}
pub(crate) fn with_residency(
inner: MetalBuffer,
dtype: DType,
shape: Vec<usize>,
residency_set: ResidencySet,
) -> Self {
residency_set.add_allocation(&inner);
let data_byte_len = inner.length() as usize;
Self {
storage: Arc::new(MlxBufferStorage {
inner,
residency_set: Some(residency_set),
file_backing: None,
cpu_writable: true,
}),
dtype,
shape,
byte_offset: 0,
data_byte_len,
}
}
pub(crate) fn from_file_mapping(
inner: MetalBuffer,
dtype: DType,
shape: Vec<usize>,
byte_offset: u64,
data_byte_len: usize,
file_backing: Arc<Mmap>,
residency_set: Option<ResidencySet>,
) -> Self {
if let Some(set) = residency_set.as_ref() {
set.add_allocation(&inner);
}
Self {
storage: Arc::new(MlxBufferStorage {
inner,
residency_set,
file_backing: Some(file_backing),
cpu_writable: false,
}),
dtype,
shape,
byte_offset,
data_byte_len,
}
}
pub(crate) fn data_view(
&self,
relative_byte_offset: usize,
data_byte_len: usize,
dtype: DType,
shape: Vec<usize>,
) -> Result<Self> {
let relative_end = relative_byte_offset
.checked_add(data_byte_len)
.ok_or_else(|| MlxError::InvalidArgument("Buffer data view range overflow".into()))?;
if relative_end > self.data_byte_len {
return Err(MlxError::InvalidArgument(format!(
"Buffer data view [{relative_byte_offset}, {relative_end}) exceeds logical data length {}",
self.data_byte_len
)));
}
let byte_offset = self
.byte_offset
.checked_add(relative_byte_offset as u64)
.ok_or_else(|| MlxError::InvalidArgument("Buffer data view offset overflow".into()))?;
let physical_end = usize::try_from(byte_offset)
.ok()
.and_then(|offset| offset.checked_add(data_byte_len))
.ok_or_else(|| MlxError::InvalidArgument("Buffer data view range overflow".into()))?;
if physical_end > self.byte_len() {
return Err(MlxError::InvalidArgument(format!(
"Buffer data view ends at {physical_end}, beyond Metal length {}",
self.byte_len()
)));
}
Ok(Self {
storage: self.storage.clone(),
dtype,
shape,
byte_offset,
data_byte_len,
})
}
#[inline]
pub fn slice_view(&self, byte_offset: u64, n_elements: usize) -> Self {
let view_byte_len = n_elements
.checked_mul(self.dtype.size_of())
.expect("slice_view: byte length overflow");
let relative_end = usize::try_from(byte_offset)
.ok()
.and_then(|offset| offset.checked_add(view_byte_len))
.expect("slice_view: range overflow");
assert!(
relative_end <= self.data_byte_len,
"slice_view: out of logical bounds (byte_offset={}, n_elements={}, dtype_size={}, data_len={})",
byte_offset,
n_elements,
self.dtype.size_of(),
self.data_byte_len
);
let absolute_byte_offset = self
.byte_offset
.checked_add(byte_offset)
.expect("slice_view: absolute offset overflow");
Self {
storage: self.storage.clone(),
dtype: self.dtype,
shape: vec![n_elements],
byte_offset: absolute_byte_offset,
data_byte_len: view_byte_len,
}
}
#[inline]
pub fn dtype(&self) -> DType {
self.dtype
}
#[inline]
pub fn shape(&self) -> &[usize] {
&self.shape
}
#[inline]
pub fn byte_len(&self) -> usize {
self.storage.inner.length() as usize
}
#[inline]
pub fn data_byte_len(&self) -> usize {
self.data_byte_len
}
#[inline]
pub fn element_count(&self) -> usize {
self.shape.iter().copied().product()
}
#[inline]
pub fn contents_ptr(&self) -> *mut std::ffi::c_void {
self.storage.inner.contents()
}
#[inline]
pub fn metal_buffer(&self) -> &MetalBuffer {
&self.storage.inner
}
#[inline]
pub fn byte_offset(&self) -> u64 {
self.byte_offset
}
#[inline]
pub fn is_file_backed(&self) -> bool {
self.storage.file_backing.is_some()
}
#[inline]
pub fn is_cpu_writable(&self) -> bool {
self.storage.cpu_writable
}
#[inline]
pub(crate) fn into_inner(self) -> MetalBuffer {
self.storage.inner.clone()
}
#[inline]
pub(crate) fn residency_set(&self) -> Option<&ResidencySet> {
self.storage.residency_set.as_ref()
}
pub fn as_slice<T: bytemuck::Pod>(&self) -> Result<&[T]> {
let elem_size = std::mem::size_of::<T>();
if elem_size == 0 {
return Err(MlxError::InvalidArgument(
"Cannot view buffer as zero-sized type".into(),
));
}
let byte_len = self.data_byte_len;
if byte_len % elem_size != 0 {
return Err(MlxError::InvalidArgument(format!(
"Buffer byte length {byte_len} is not a multiple of element size {elem_size}"
)));
}
let base = self.contents_ptr();
if base.is_null() {
return Err(MlxError::BufferAllocationError { bytes: byte_len });
}
let ptr = unsafe { (base as *const u8).add(self.byte_offset as usize) };
if (ptr as usize) % std::mem::align_of::<T>() != 0 {
return Err(MlxError::InvalidArgument(format!(
"Buffer data offset {} is not aligned for {}",
self.byte_offset,
std::any::type_name::<T>()
)));
}
let count = byte_len / elem_size;
let slice = unsafe { std::slice::from_raw_parts(ptr as *const T, count) };
Ok(slice)
}
pub fn as_mut_slice<T: bytemuck::Pod>(&mut self) -> Result<&mut [T]> {
if !self.storage.cpu_writable {
return Err(MlxError::InvalidArgument(
"Cannot mutate a read-only file-backed buffer".into(),
));
}
let elem_size = std::mem::size_of::<T>();
if elem_size == 0 {
return Err(MlxError::InvalidArgument(
"Cannot view buffer as zero-sized type".into(),
));
}
let byte_len = self.data_byte_len;
if byte_len % elem_size != 0 {
return Err(MlxError::InvalidArgument(format!(
"Buffer byte length {byte_len} is not a multiple of element size {elem_size}"
)));
}
let base = self.contents_ptr();
if base.is_null() {
return Err(MlxError::BufferAllocationError { bytes: byte_len });
}
let ptr = unsafe { (base as *mut u8).add(self.byte_offset as usize) };
if (ptr as usize) % std::mem::align_of::<T>() != 0 {
return Err(MlxError::InvalidArgument(format!(
"Buffer data offset {} is not aligned for {}",
self.byte_offset,
std::any::type_name::<T>()
)));
}
let count = byte_len / elem_size;
let slice = unsafe { std::slice::from_raw_parts_mut(ptr as *mut T, count) };
Ok(slice)
}
pub fn as_logical_slice<T: bytemuck::Pod>(&self) -> Result<&[T]> {
let (ptr, count) = self.logical_cpu_view::<T>()?;
Ok(unsafe { std::slice::from_raw_parts(ptr as *const T, count) })
}
pub fn as_logical_mut_slice<T: bytemuck::Pod>(&mut self) -> Result<&mut [T]> {
if !self.storage.cpu_writable {
return Err(MlxError::InvalidArgument(
"Cannot mutate a read-only file-backed buffer".into(),
));
}
let (ptr, count) = self.logical_cpu_view::<T>()?;
Ok(unsafe { std::slice::from_raw_parts_mut(ptr as *mut T, count) })
}
fn logical_cpu_view<T: bytemuck::Pod>(&self) -> Result<(*mut u8, usize)> {
let elem_size = std::mem::size_of::<T>();
if elem_size == 0 {
return Err(MlxError::InvalidArgument(
"Cannot view buffer as zero-sized type".into(),
));
}
let logical_bytes = self
.element_count()
.checked_mul(self.dtype.size_of())
.ok_or_else(|| {
MlxError::InvalidArgument("Buffer logical byte length overflow".into())
})?;
if logical_bytes % elem_size != 0 {
return Err(MlxError::InvalidArgument(format!(
"Buffer logical byte length {logical_bytes} is not a multiple of element size {elem_size}"
)));
}
if logical_bytes > self.data_byte_len {
return Err(MlxError::InvalidArgument(format!(
"Buffer logical byte length {logical_bytes} exceeds data length {}",
self.data_byte_len
)));
}
let offset = self.byte_offset as usize;
let end = offset.checked_add(logical_bytes).ok_or_else(|| {
MlxError::InvalidArgument("Buffer logical CPU view range overflow".into())
})?;
let allocation_len = self.byte_len();
if end > allocation_len {
return Err(MlxError::InvalidArgument(format!(
"Buffer logical CPU view [{offset}, {end}) exceeds allocation length {allocation_len}"
)));
}
let base = self.contents_ptr();
if base.is_null() {
return Err(MlxError::BufferAllocationError {
bytes: logical_bytes,
});
}
let ptr = unsafe { (base as *mut u8).add(offset) };
if (ptr as usize) % std::mem::align_of::<T>() != 0 {
return Err(MlxError::InvalidArgument(format!(
"Buffer logical CPU view offset {offset} is not aligned for {}",
std::any::type_name::<T>()
)));
}
Ok((ptr, logical_bytes / elem_size))
}
#[allow(dead_code)]
pub(crate) fn reshape(&mut self, dtype: DType, shape: Vec<usize>) {
self.dtype = dtype;
self.shape = shape;
}
pub fn with_shape(&self, shape: Vec<usize>) -> std::result::Result<Self, MlxError> {
let numel: usize = shape.iter().product();
if numel != self.element_count() {
return Err(MlxError::InvalidArgument(format!(
"with_shape: numel({:?}) = {numel} != element_count = {}",
shape,
self.element_count(),
)));
}
Ok(Self {
storage: self.storage.clone(),
dtype: self.dtype,
shape,
byte_offset: self.byte_offset,
data_byte_len: self.data_byte_len,
})
}
}
impl fmt::Debug for MlxBuffer {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("MlxBuffer")
.field("dtype", &self.dtype)
.field("shape", &self.shape)
.field("byte_len", &self.byte_len())
.field("data_byte_len", &self.data_byte_len())
.field("byte_offset", &self.byte_offset)
.field("file_backed", &self.is_file_backed())
.finish()
}
}