use std::ffi::c_void;
use crate::{
error::{self, Result},
utils::{guard::Guarded, runtime_lock, SUCCESS},
Array, Dtype, Event, Stream,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum HostTransferPolicy {
Transfer,
Managed,
}
impl HostTransferPolicy {
fn as_raw(self) -> safemlx_sys::mlx_host_transfer_policy {
match self {
Self::Transfer => {
safemlx_sys::mlx_host_transfer_policy__MLX_HOST_TRANSFER_POLICY_TRANSFER
}
Self::Managed => {
safemlx_sys::mlx_host_transfer_policy__MLX_HOST_TRANSFER_POLICY_MANAGED
}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum HostTransferStorageKind {
Cpu,
MetalShared,
CudaPinned,
CudaManaged,
}
impl HostTransferStorageKind {
fn as_raw(self) -> safemlx_sys::mlx_host_transfer_storage_kind {
match self {
Self::Cpu => safemlx_sys::mlx_host_transfer_storage_kind__MLX_HOST_TRANSFER_STORAGE_CPU,
Self::MetalShared => {
safemlx_sys::mlx_host_transfer_storage_kind__MLX_HOST_TRANSFER_STORAGE_METAL_SHARED
}
Self::CudaPinned => {
safemlx_sys::mlx_host_transfer_storage_kind__MLX_HOST_TRANSFER_STORAGE_CUDA_PINNED
}
Self::CudaManaged => {
safemlx_sys::mlx_host_transfer_storage_kind__MLX_HOST_TRANSFER_STORAGE_CUDA_MANAGED
}
}
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
pub struct HostTransferMemoryStats {
pub active_bytes: usize,
pub peak_bytes: usize,
pub active_allocations: usize,
pub peak_allocations: usize,
}
fn check_status(status: i32) -> Result<()> {
if status == SUCCESS {
Ok(())
} else {
Err(error::get_and_clear_last_mlx_error()
.expect("MLX host-transfer operation failed but no error was set")
.into())
}
}
pub fn host_transfer_memory_stats(
kind: HostTransferStorageKind,
) -> Result<HostTransferMemoryStats> {
let _guard = runtime_lock::enter();
error::ensure_mlx_error_handler();
let mut stats = safemlx_sys::mlx_host_transfer_memory_stats {
active_bytes: 0,
peak_bytes: 0,
active_allocations: 0,
peak_allocations: 0,
};
check_status(unsafe {
safemlx_sys::mlx_host_transfer_memory_stats_get(&mut stats, kind.as_raw())
})?;
Ok(HostTransferMemoryStats {
active_bytes: stats.active_bytes,
peak_bytes: stats.peak_bytes,
active_allocations: stats.active_allocations,
peak_allocations: stats.peak_allocations,
})
}
pub fn reset_host_transfer_peak_memory(kind: HostTransferStorageKind) -> Result<()> {
let _guard = runtime_lock::enter();
error::ensure_mlx_error_handler();
check_status(unsafe { safemlx_sys::mlx_host_transfer_memory_stats_reset_peak(kind.as_raw()) })
}
pub fn host_transfer_capacity_upper_bound(
nbytes: usize,
policy: HostTransferPolicy,
) -> Result<usize> {
let _guard = runtime_lock::enter();
error::ensure_mlx_error_handler();
let mut capacity = 0;
check_status(unsafe {
safemlx_sys::mlx_host_transfer_capacity_upper_bound(&mut capacity, nbytes, policy.as_raw())
})?;
Ok(capacity)
}
pub struct HostTransferBuffer {
pub(crate) raw: safemlx_sys::mlx_host_transfer_buffer,
}
impl HostTransferBuffer {
#[track_caller]
pub fn new(shape: &[i32], dtype: Dtype, policy: HostTransferPolicy) -> Result<Self> {
let _guard = runtime_lock::enter();
let dim = i32::try_from(shape.len()).map_err(|_| crate::error::Exception {
what: "Host transfer buffer rank exceeds i32::MAX".to_string(),
location: std::panic::Location::caller(),
})?;
Self::try_from_op(|buffer| unsafe {
safemlx_sys::mlx_host_transfer_buffer_new(
buffer,
shape.as_ptr(),
dim,
dtype.into(),
policy.as_raw(),
)
})
}
pub fn copy_from_array(
source: &Array,
policy: HostTransferPolicy,
stream: impl AsRef<Stream>,
) -> Result<PendingHostTransfer> {
let _guard = runtime_lock::enter();
let (buffer, completion) = <(Self, Event)>::try_from_op(|(buffer, event)| unsafe {
safemlx_sys::mlx_copy_to_host(
buffer,
event,
source.as_ptr(),
policy.as_raw(),
stream.as_ref().as_ptr(),
)
})?;
Ok(PendingHostTransfer { buffer, completion })
}
pub fn copy_to_array(self, stream: impl AsRef<Stream>) -> Result<PendingDeviceTransfer> {
let _guard = runtime_lock::enter();
let (value, completion) = <(Array, Event)>::try_from_op(|(array, event)| unsafe {
safemlx_sys::mlx_copy_from_host(array, event, self.raw, stream.as_ref().as_ptr())
})?;
Ok(PendingDeviceTransfer {
source: self,
value,
completion,
})
}
pub fn freeze(self) -> ImmutableHostTransferBuffer {
ImmutableHostTransferBuffer { buffer: self }
}
pub fn shape(&self) -> Result<Vec<i32>> {
let _guard = runtime_lock::enter();
let ndim = usize::try_from_op(|output| unsafe {
safemlx_sys::mlx_host_transfer_buffer_ndim(output, self.raw)
})?;
let mut shape = std::ptr::null();
let status = unsafe { safemlx_sys::mlx_host_transfer_buffer_shape(&mut shape, self.raw) };
if status != SUCCESS {
return <() as Guarded>::try_from_op(|_| status).map(|_| Vec::new());
}
if ndim == 0 {
return Ok(Vec::new());
}
debug_assert!(!shape.is_null());
Ok(unsafe { std::slice::from_raw_parts(shape, ndim) }.to_vec())
}
pub fn dtype(&self) -> Result<Dtype> {
let _guard = runtime_lock::enter();
let raw = u32::try_from_op(|output| unsafe {
safemlx_sys::mlx_host_transfer_buffer_dtype(output.cast(), self.raw)
})?;
Ok(Dtype::try_from(raw).expect("MLX returned an unknown dtype"))
}
pub fn len(&self) -> Result<usize> {
let _guard = runtime_lock::enter();
usize::try_from_op(|output| unsafe {
safemlx_sys::mlx_host_transfer_buffer_size(output, self.raw)
})
}
pub fn is_empty(&self) -> Result<bool> {
self.len().map(|len| len == 0)
}
pub fn nbytes(&self) -> Result<usize> {
let _guard = runtime_lock::enter();
usize::try_from_op(|output| unsafe {
safemlx_sys::mlx_host_transfer_buffer_nbytes(output, self.raw)
})
}
pub fn capacity(&self) -> Result<usize> {
let _guard = runtime_lock::enter();
usize::try_from_op(|output| unsafe {
safemlx_sys::mlx_host_transfer_buffer_capacity(output, self.raw)
})
}
pub fn policy(&self) -> Result<HostTransferPolicy> {
let _guard = runtime_lock::enter();
let raw = u32::try_from_op(|output| unsafe {
safemlx_sys::mlx_host_transfer_buffer_policy(output.cast(), self.raw)
})?;
match raw {
safemlx_sys::mlx_host_transfer_policy__MLX_HOST_TRANSFER_POLICY_TRANSFER => {
Ok(HostTransferPolicy::Transfer)
}
safemlx_sys::mlx_host_transfer_policy__MLX_HOST_TRANSFER_POLICY_MANAGED => {
Ok(HostTransferPolicy::Managed)
}
_ => unreachable!("MLX returned an unknown host transfer policy"),
}
}
pub fn storage_kind(&self) -> Result<HostTransferStorageKind> {
let _guard = runtime_lock::enter();
let raw = u32::try_from_op(|output| unsafe {
safemlx_sys::mlx_host_transfer_buffer_storage_kind(output.cast(), self.raw)
})?;
match raw {
safemlx_sys::mlx_host_transfer_storage_kind__MLX_HOST_TRANSFER_STORAGE_CPU => {
Ok(HostTransferStorageKind::Cpu)
}
safemlx_sys::mlx_host_transfer_storage_kind__MLX_HOST_TRANSFER_STORAGE_METAL_SHARED => {
Ok(HostTransferStorageKind::MetalShared)
}
safemlx_sys::mlx_host_transfer_storage_kind__MLX_HOST_TRANSFER_STORAGE_CUDA_PINNED => {
Ok(HostTransferStorageKind::CudaPinned)
}
safemlx_sys::mlx_host_transfer_storage_kind__MLX_HOST_TRANSFER_STORAGE_CUDA_MANAGED => {
Ok(HostTransferStorageKind::CudaManaged)
}
_ => unreachable!("MLX returned an unknown host transfer storage kind"),
}
}
pub fn as_bytes(&self) -> Result<&[u8]> {
let _guard = runtime_lock::enter();
let len = self.nbytes()?;
let mut pointer: *const c_void = std::ptr::null();
let status = unsafe { safemlx_sys::mlx_host_transfer_buffer_data(&mut pointer, self.raw) };
if status != SUCCESS {
return <() as Guarded>::try_from_op(|_| status).map(|_| &[][..]);
}
if len == 0 {
return Ok(&[]);
}
debug_assert!(!pointer.is_null());
Ok(unsafe { std::slice::from_raw_parts(pointer.cast(), len) })
}
pub fn as_bytes_mut(&mut self) -> Result<&mut [u8]> {
let _guard = runtime_lock::enter();
let len = self.nbytes()?;
let mut pointer: *mut c_void = std::ptr::null_mut();
let status =
unsafe { safemlx_sys::mlx_host_transfer_buffer_data_mut(&mut pointer, self.raw) };
if status != SUCCESS {
return <() as Guarded>::try_from_op(|_| status).map(|_| &mut [][..]);
}
if len == 0 {
return Ok(&mut []);
}
debug_assert!(!pointer.is_null());
Ok(unsafe { std::slice::from_raw_parts_mut(pointer.cast(), len) })
}
}
impl Drop for HostTransferBuffer {
fn drop(&mut self) {
let _guard = runtime_lock::enter();
let status = unsafe { safemlx_sys::mlx_host_transfer_buffer_free(self.raw) };
debug_assert_eq!(status, SUCCESS);
}
}
impl std::fmt::Debug for HostTransferBuffer {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("HostTransferBuffer")
.field("shape", &self.shape())
.field("dtype", &self.dtype())
.field("nbytes", &self.nbytes())
.field("capacity", &self.capacity())
.field("policy", &self.policy())
.field("storage_kind", &self.storage_kind())
.finish()
}
}
pub struct ImmutableHostTransferBuffer {
buffer: HostTransferBuffer,
}
unsafe impl Send for ImmutableHostTransferBuffer {}
unsafe impl Sync for ImmutableHostTransferBuffer {}
impl ImmutableHostTransferBuffer {
pub fn shape(&self) -> Result<Vec<i32>> {
self.buffer.shape()
}
pub fn dtype(&self) -> Result<Dtype> {
self.buffer.dtype()
}
pub fn len(&self) -> Result<usize> {
self.buffer.len()
}
pub fn is_empty(&self) -> Result<bool> {
self.buffer.is_empty()
}
pub fn nbytes(&self) -> Result<usize> {
self.buffer.nbytes()
}
pub fn capacity(&self) -> Result<usize> {
self.buffer.capacity()
}
pub fn policy(&self) -> Result<HostTransferPolicy> {
self.buffer.policy()
}
pub fn storage_kind(&self) -> Result<HostTransferStorageKind> {
self.buffer.storage_kind()
}
pub fn as_bytes(&self) -> Result<&[u8]> {
self.buffer.as_bytes()
}
pub fn copy_to_array(&self, stream: impl AsRef<Stream>) -> Result<SubmittedDeviceTransfer> {
let _guard = runtime_lock::enter();
let (value, completion) = <(Array, Event)>::try_from_op(|(array, event)| unsafe {
safemlx_sys::mlx_copy_from_host(array, event, self.buffer.raw, stream.as_ref().as_ptr())
})?;
Ok(SubmittedDeviceTransfer { value, completion })
}
}
impl std::fmt::Debug for ImmutableHostTransferBuffer {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("ImmutableHostTransferBuffer")
.field("shape", &self.shape())
.field("dtype", &self.dtype())
.field("nbytes", &self.nbytes())
.field("capacity", &self.capacity())
.field("policy", &self.policy())
.field("storage_kind", &self.storage_kind())
.finish()
}
}
#[derive(Debug)]
pub struct SubmittedDeviceTransfer {
value: Array,
completion: Event,
}
impl SubmittedDeviceTransfer {
pub fn value(&self) -> &Array {
&self.value
}
pub fn completion(&self) -> &Event {
&self.completion
}
pub fn into_parts(self) -> (Array, Event) {
(self.value, self.completion)
}
pub fn synchronize(self) -> Result<Array> {
self.completion.synchronize()?;
Ok(self.value)
}
}
#[derive(Debug)]
pub struct PendingHostTransfer {
buffer: HostTransferBuffer,
completion: Event,
}
impl PendingHostTransfer {
pub fn completion(&self) -> &Event {
&self.completion
}
pub fn into_parts(self) -> (HostTransferBuffer, Event) {
(self.buffer, self.completion)
}
pub fn synchronize(self) -> Result<HostTransferBuffer> {
self.completion.synchronize()?;
Ok(self.buffer)
}
}
#[derive(Debug)]
pub struct PendingDeviceTransfer {
source: HostTransferBuffer,
value: Array,
completion: Event,
}
impl PendingDeviceTransfer {
pub fn completion(&self) -> &Event {
&self.completion
}
pub fn synchronize(self) -> Result<(Array, HostTransferBuffer)> {
self.completion.synchronize()?;
Ok((self.value, self.source))
}
}