From a2ac7ffcd9aba2455693a7005356c6dd0afb66d1 Mon Sep 17 00:00:00 2001 From: kerthcet Date: Mon, 20 Jul 2026 16:58:15 +0100 Subject: [PATCH 1/9] support stream Signed-off-by: kerthcet --- crates/mlx/examples/hello.rs | 16 ++- crates/mlx/src/array.rs | 269 ++++++++++++++++++++++++++++++++++- crates/mlx/src/dtype.rs | 54 +++++-- crates/mlx/src/lib.rs | 2 + crates/mlx/src/stream.rs | 94 ++++++++++++ 5 files changed, 418 insertions(+), 17 deletions(-) create mode 100644 crates/mlx/src/stream.rs diff --git a/crates/mlx/examples/hello.rs b/crates/mlx/examples/hello.rs index a74a468..de62b83 100644 --- a/crates/mlx/examples/hello.rs +++ b/crates/mlx/examples/hello.rs @@ -5,7 +5,7 @@ //! cargo run --example hello //! ``` -use mlx::Array; +use mlx::{Array, Stream}; fn main() { println!("MLX version: {}", mlx::version()); @@ -15,4 +15,18 @@ fn main() { println!("size: {}", a.size()); println!("ndim: {}", a.ndim()); println!("array:\n{a:?}"); + + let b = Array::from_slice(&[6.0f32, 5.0, 4.0, 3.0, 2.0, 1.0], &[2, 3]); + + // Operators use the default stream. + println!("a + b:\n{:?}", &a + &b); + + // Or steer the default onto a specific device once, up front. + Stream::cpu().set_as_default(); + let total = a.sum(false, &Stream::default()); + println!("sum(a) on CPU: {}", total.item::()); + + // Explicit stream control is still available per-op. + let product = a.multiply(&b, &Stream::gpu()); + println!("a * b (GPU):\n{product:?}"); } diff --git a/crates/mlx/src/array.rs b/crates/mlx/src/array.rs index d09b916..4b9eedd 100644 --- a/crates/mlx/src/array.rs +++ b/crates/mlx/src/array.rs @@ -6,6 +6,7 @@ use std::fmt; use mlx_sys as sys; use crate::dtype::ArrayElement; +use crate::stream::Stream; /// An N-dimensional MLX array. /// @@ -24,8 +25,6 @@ impl Array { } /// Returns the raw handle. The `Array` retains ownership. - // Scaffolding for ops that need the underlying handle; unused for now. - #[allow(dead_code)] pub(crate) fn as_raw(&self) -> sys::mlx_array { self.handle } @@ -76,6 +75,155 @@ impl Array { let ptr = unsafe { sys::mlx_array_shape(self.handle) }; (0..ndim).map(|i| unsafe { *ptr.add(i) }).collect() } + + /// Forces evaluation of this array. + /// + /// MLX is lazy: ops build a graph and only compute when the result is + /// needed. `eval` materializes the values now. + pub fn eval(&self) { + // SAFETY: handle is valid for the lifetime of `self`. + unsafe { + sys::mlx_array_eval(self.handle); + } + } + + /// Reads the value of a scalar (single-element) array. + /// + /// The element type `T` selects the accessor at compile time, e.g. + /// `a.item::()`. Evaluates the array first. MLX casts the stored dtype + /// to `T`. + pub fn item(&self) -> T { + self.eval(); + // SAFETY: the array is evaluated above; `read_item` picks the accessor + // matching `T`. + unsafe { T::read_item(self.handle) } + } + + /// Copies the array's contents into a `Vec`, row-major. + /// + /// The element type `T` selects the accessor at compile time, e.g. + /// `a.to_vec::()`. Evaluates the array first. + /// + /// # Panics + /// Panics if `T::DTYPE` does not match the array's dtype. + pub fn to_vec(&self) -> Vec { + self.eval(); + // SAFETY: mlx_array_dtype reads a valid handle. + let dtype = unsafe { sys::mlx_array_dtype(self.handle) }; + assert_eq!( + dtype, + T::DTYPE, + "array dtype does not match requested element type" + ); + let len = self.size(); + // SAFETY: dtype matches `T` (checked above), so the pointer is valid for + // `size` contiguous `T` until the array is mutated or freed. + let ptr = unsafe { T::data_ptr(self.handle) }; + (0..len).map(|i| unsafe { *ptr.add(i) }).collect() + } + + /// Elementwise addition: `self + other`. + pub fn add(&self, other: &Array, stream: &Stream) -> Array { + self.binary_op(other, stream, sys::mlx_add) + } + + /// Elementwise subtraction: `self - other`. + pub fn subtract(&self, other: &Array, stream: &Stream) -> Array { + self.binary_op(other, stream, sys::mlx_subtract) + } + + /// Elementwise multiplication: `self * other`. + pub fn multiply(&self, other: &Array, stream: &Stream) -> Array { + self.binary_op(other, stream, sys::mlx_multiply) + } + + /// Elementwise division: `self / other`. + pub fn divide(&self, other: &Array, stream: &Stream) -> Array { + self.binary_op(other, stream, sys::mlx_divide) + } + + /// Elementwise square root. + pub fn sqrt(&self, stream: &Stream) -> Array { + self.unary_op(stream, sys::mlx_sqrt) + } + + /// Elementwise exponential. + pub fn exp(&self, stream: &Stream) -> Array { + self.unary_op(stream, sys::mlx_exp) + } + + /// Elementwise absolute value. + pub fn abs(&self, stream: &Stream) -> Array { + self.unary_op(stream, sys::mlx_abs) + } + + /// Elementwise negation. + pub fn negative(&self, stream: &Stream) -> Array { + self.unary_op(stream, sys::mlx_negative) + } + + /// Sum of all elements, returning a scalar array. + /// + /// With `keepdims == false` the result is 0-dimensional. + pub fn sum(&self, keepdims: bool, stream: &Stream) -> Array { + self.reduce_op(keepdims, stream, sys::mlx_sum) + } + + /// Mean of all elements, returning a scalar array. + /// + /// With `keepdims == false` the result is 0-dimensional. + pub fn mean(&self, keepdims: bool, stream: &Stream) -> Array { + self.reduce_op(keepdims, stream, sys::mlx_mean) + } + + /// Shared plumbing for `res = op(a, b, stream)` binary ops. + fn binary_op( + &self, + other: &Array, + stream: &Stream, + op: unsafe extern "C" fn( + *mut sys::mlx_array, + sys::mlx_array, + sys::mlx_array, + sys::mlx_stream, + ) -> i32, + ) -> Array { + let mut out = unsafe { sys::mlx_array_new() }; + // SAFETY: all handles are valid; `op` writes the result into `out`. + unsafe { + op(&mut out, self.handle, other.as_raw(), stream.as_raw()); + Self::from_raw(out) + } + } + + /// Shared plumbing for `res = op(a, stream)` unary ops. + fn unary_op( + &self, + stream: &Stream, + op: unsafe extern "C" fn(*mut sys::mlx_array, sys::mlx_array, sys::mlx_stream) -> i32, + ) -> Array { + let mut out = unsafe { sys::mlx_array_new() }; + // SAFETY: handle/stream are valid; `op` writes the result into `out`. + unsafe { + op(&mut out, self.handle, stream.as_raw()); + Self::from_raw(out) + } + } + + /// Shared plumbing for `res = op(a, keepdims, stream)` full reductions. + fn reduce_op( + &self, + keepdims: bool, + stream: &Stream, + op: unsafe extern "C" fn(*mut sys::mlx_array, sys::mlx_array, bool, sys::mlx_stream) -> i32, + ) -> Array { + let mut out = unsafe { sys::mlx_array_new() }; + // SAFETY: handle/stream are valid; `op` writes the result into `out`. + unsafe { + op(&mut out, self.handle, keepdims, stream.as_raw()); + Self::from_raw(out) + } + } } impl fmt::Debug for Array { @@ -99,6 +247,39 @@ impl Drop for Array { } } +// Arithmetic operators run on the current default stream (see +// [`Stream::set_as_default`]). For explicit stream control, call the inherent +// methods (`a.add(&b, &stream)`) instead. +// +// Implemented on `&Array` so operands are borrowed, not consumed: `&a + &b` +// leaves both arrays usable afterwards. +macro_rules! impl_binop { + ($($trait:ident :: $method:ident => $op:ident),* $(,)?) => { + $( + impl std::ops::$trait for &Array { + type Output = Array; + fn $method(self, rhs: &Array) -> Array { + self.$op(rhs, &Stream::default()) + } + } + )* + }; +} + +impl_binop! { + Add::add => add, + Sub::sub => subtract, + Mul::mul => multiply, + Div::div => divide, +} + +impl std::ops::Neg for &Array { + type Output = Array; + fn neg(self) -> Array { + self.negative(&Stream::default()) + } +} + #[cfg(test)] mod tests { use super::*; @@ -141,4 +322,88 @@ mod tests { let s = format!("{a:?}"); assert!(s.contains("array"), "unexpected debug output: {s}"); } + + #[test] + fn binary_ops_compute_elementwise() { + let s = Stream::cpu(); + let a = Array::from_slice(&[1.0f32, 2.0, 3.0], &[3]); + let b = Array::from_slice(&[4.0f32, 5.0, 6.0], &[3]); + + assert_eq!(a.add(&b, &s).to_vec::(), vec![5.0, 7.0, 9.0]); + assert_eq!(b.subtract(&a, &s).to_vec::(), vec![3.0, 3.0, 3.0]); + assert_eq!(a.multiply(&b, &s).to_vec::(), vec![4.0, 10.0, 18.0]); + assert_eq!(b.divide(&a, &s).to_vec::(), vec![4.0, 2.5, 2.0]); + } + + #[test] + fn unary_ops_compute_elementwise() { + let s = Stream::cpu(); + let a = Array::from_slice(&[1.0f32, 4.0, 9.0], &[3]); + assert_eq!(a.sqrt(&s).to_vec::(), vec![1.0, 2.0, 3.0]); + + let b = Array::from_slice(&[-1.0f32, 2.0, -3.0], &[3]); + assert_eq!(b.abs(&s).to_vec::(), vec![1.0, 2.0, 3.0]); + assert_eq!(b.negative(&s).to_vec::(), vec![1.0, -2.0, 3.0]); + } + + #[test] + fn reductions_produce_scalars() { + let s = Stream::cpu(); + let a = Array::from_slice(&[1.0f32, 2.0, 3.0, 4.0], &[4]); + + let sum = a.sum(false, &s); + assert_eq!(sum.ndim(), 0); + assert_eq!(sum.item::(), 10.0); + + assert_eq!(a.mean(false, &s).item::(), 2.5); + } + + #[test] + fn keepdims_retains_rank() { + let s = Stream::cpu(); + let a = Array::from_slice(&[1.0f32, 2.0, 3.0, 4.0], &[2, 2]); + let sum = a.sum(true, &s); + assert_eq!(sum.shape(), vec![1, 1]); + assert_eq!(sum.item::(), 10.0); + } + + #[test] + fn item_reads_scalar() { + let a = Array::from_slice(&[42.0f32], &[]); + assert_eq!(a.item::(), 42.0); + let b = Array::from_slice(&[7i32], &[]); + assert_eq!(b.item::(), 7); + } + + #[test] + fn to_vec_is_generic_over_dtype() { + let ints = Array::from_slice(&[1i32, 2, 3], &[3]); + assert_eq!(ints.to_vec::(), vec![1, 2, 3]); + let floats = Array::from_slice(&[1.5f32, 2.5], &[2]); + assert_eq!(floats.to_vec::(), vec![1.5, 2.5]); + } + + #[test] + #[should_panic(expected = "does not match requested element type")] + fn to_vec_wrong_dtype_panics() { + let ints = Array::from_slice(&[1i32, 2, 3], &[3]); + let _ = ints.to_vec::(); + } + + #[test] + fn operators_match_methods() { + // Operators run on the default stream; results should equal the + // explicit-method equivalents. + let a = Array::from_slice(&[10.0f32, 20.0, 30.0], &[3]); + let b = Array::from_slice(&[1.0f32, 2.0, 3.0], &[3]); + + assert_eq!((&a + &b).to_vec::(), vec![11.0, 22.0, 33.0]); + assert_eq!((&a - &b).to_vec::(), vec![9.0, 18.0, 27.0]); + assert_eq!((&a * &b).to_vec::(), vec![10.0, 40.0, 90.0]); + assert_eq!((&a / &b).to_vec::(), vec![10.0, 10.0, 10.0]); + assert_eq!((-&a).to_vec::(), vec![-10.0, -20.0, -30.0]); + + // Operands are borrowed, so `a` is still usable here. + assert_eq!(a.to_vec::(), vec![10.0, 20.0, 30.0]); + } } diff --git a/crates/mlx/src/dtype.rs b/crates/mlx/src/dtype.rs index c359c4f..63b4776 100644 --- a/crates/mlx/src/dtype.rs +++ b/crates/mlx/src/dtype.rs @@ -9,18 +9,44 @@ mod sealed { /// A Rust type that has a corresponding MLX [`mlx_dtype`](sys::mlx_dtype). /// /// This trait is sealed: it can only be implemented for the primitive types -/// MLX supports, so `T::DTYPE` is always a valid dtype. -pub trait ArrayElement: sealed::Sealed + Copy { +/// MLX supports, so `T::DTYPE` is always a valid dtype and the accessors below +/// always match it. +pub trait ArrayElement: sealed::Sealed + Copy + Default { /// The MLX dtype corresponding to this Rust type. const DTYPE: sys::mlx_dtype; + + /// Reads the value of a scalar (single-element) array as this type. + /// + /// # Safety + /// `arr` must be a valid, already-evaluated scalar `mlx_array`. + unsafe fn read_item(arr: sys::mlx_array) -> Self; + + /// Returns a pointer to this array's contiguous data of this type. + /// + /// # Safety + /// `arr` must be a valid, already-evaluated `mlx_array` whose dtype is + /// `Self::DTYPE`. The pointer is valid until `arr` is mutated or freed. + unsafe fn data_ptr(arr: sys::mlx_array) -> *const Self; } macro_rules! impl_array_element { - ($($rust:ty => $dtype:expr),* $(,)?) => { + ($($rust:ty => $dtype:expr, $item:path, $data:path),* $(,)?) => { $( impl sealed::Sealed for $rust {} impl ArrayElement for $rust { const DTYPE: sys::mlx_dtype = $dtype; + + unsafe fn read_item(arr: sys::mlx_array) -> Self { + let mut out = <$rust>::default(); + // SAFETY: caller guarantees `arr` is a valid scalar array. + unsafe { $item(&mut out, arr); } + out + } + + unsafe fn data_ptr(arr: sys::mlx_array) -> *const Self { + // SAFETY: caller guarantees `arr` is valid with dtype DTYPE. + unsafe { $data(arr) } + } } )* }; @@ -30,17 +56,17 @@ macro_rules! impl_array_element { // without a stable Rust equivalent (float16, bfloat16, complex64) are left for // dedicated newtypes later. impl_array_element! { - bool => sys::MLX_BOOL, - u8 => sys::MLX_UINT8, - u16 => sys::MLX_UINT16, - u32 => sys::MLX_UINT32, - u64 => sys::MLX_UINT64, - i8 => sys::MLX_INT8, - i16 => sys::MLX_INT16, - i32 => sys::MLX_INT32, - i64 => sys::MLX_INT64, - f32 => sys::MLX_FLOAT32, - f64 => sys::MLX_FLOAT64, + bool => sys::MLX_BOOL, sys::mlx_array_item_bool, sys::mlx_array_data_bool, + u8 => sys::MLX_UINT8, sys::mlx_array_item_uint8, sys::mlx_array_data_uint8, + u16 => sys::MLX_UINT16, sys::mlx_array_item_uint16, sys::mlx_array_data_uint16, + u32 => sys::MLX_UINT32, sys::mlx_array_item_uint32, sys::mlx_array_data_uint32, + u64 => sys::MLX_UINT64, sys::mlx_array_item_uint64, sys::mlx_array_data_uint64, + i8 => sys::MLX_INT8, sys::mlx_array_item_int8, sys::mlx_array_data_int8, + i16 => sys::MLX_INT16, sys::mlx_array_item_int16, sys::mlx_array_data_int16, + i32 => sys::MLX_INT32, sys::mlx_array_item_int32, sys::mlx_array_data_int32, + i64 => sys::MLX_INT64, sys::mlx_array_item_int64, sys::mlx_array_data_int64, + f32 => sys::MLX_FLOAT32, sys::mlx_array_item_float32, sys::mlx_array_data_float32, + f64 => sys::MLX_FLOAT64, sys::mlx_array_item_float64, sys::mlx_array_data_float64, } #[cfg(test)] diff --git a/crates/mlx/src/lib.rs b/crates/mlx/src/lib.rs index 1b22be7..1c4812b 100644 --- a/crates/mlx/src/lib.rs +++ b/crates/mlx/src/lib.rs @@ -5,9 +5,11 @@ mod array; mod dtype; +mod stream; pub use array::Array; pub use dtype::ArrayElement; +pub use stream::Stream; /// Returns the version string of the underlying MLX library. pub fn version() -> String { diff --git a/crates/mlx/src/stream.rs b/crates/mlx/src/stream.rs new file mode 100644 index 0000000..df628cd --- /dev/null +++ b/crates/mlx/src/stream.rs @@ -0,0 +1,94 @@ +//! A safe wrapper around `mlx_stream`. +//! +//! MLX schedules every operation on a [`Stream`], which is bound to a device +//! (CPU or GPU). Most ops take a stream argument. MLX's own default is the GPU +//! when a GPU backend is available (as on Apple Silicon), else the CPU. + +use mlx_sys as sys; + +/// An execution stream bound to a device. +pub struct Stream { + handle: sys::mlx_stream, +} + +impl Stream { + /// The default stream on the GPU device. + pub fn gpu() -> Self { + // SAFETY: constructor returns an owned handle. + let handle = unsafe { sys::mlx_default_gpu_stream_new() }; + Self { handle } + } + + /// The default stream on the CPU device. + pub fn cpu() -> Self { + // SAFETY: constructor returns an owned handle. + let handle = unsafe { sys::mlx_default_cpu_stream_new() }; + Self { handle } + } + + /// Returns the raw handle. The `Stream` retains ownership. + pub(crate) fn as_raw(&self) -> sys::mlx_stream { + self.handle + } + + /// Makes this the process-wide default stream (and its device the default + /// device). + /// + /// Operations that don't take an explicit stream — the arithmetic operators + /// (`&a + &b`) and anything built on [`Stream::default`] — resolve the + /// current default. Call this once at startup to steer them onto a chosen + /// device. + pub fn set_as_default(&self) { + // SAFETY: `handle` is a valid stream owned by `self`; mlx copies out of + // the pointers we pass and out of the handle. + unsafe { + let mut dev = sys::mlx_device_new(); + sys::mlx_stream_get_device(&mut dev, self.handle); + sys::mlx_set_default_device(dev); + sys::mlx_device_free(dev); + sys::mlx_set_default_stream(self.handle); + } + } +} + +impl Default for Stream { + /// MLX's *current* default stream — follows [`Stream::set_as_default`]. + /// + /// The initial default is chosen by MLX, not this crate: it is the GPU when + /// a GPU backend (Metal) is available, otherwise the CPU. This only reads + /// that default; use [`Stream::set_as_default`] to change it. + fn default() -> Self { + // SAFETY: out-params are valid; mlx writes owned handles into them. + let handle = unsafe { + let mut dev = sys::mlx_device_new(); + sys::mlx_get_default_device(&mut dev); + let mut stream = sys::mlx_stream_new(); + sys::mlx_get_default_stream(&mut stream, dev); + sys::mlx_device_free(dev); + stream + }; + Self { handle } + } +} + +impl Drop for Stream { + fn drop(&mut self) { + // SAFETY: `handle` was created by mlx and is owned solely by `self`. + unsafe { + sys::mlx_stream_free(self.handle); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn constructs_default_streams() { + // Just exercises construction + drop on both devices. + let _gpu = Stream::gpu(); + let _cpu = Stream::cpu(); + let _default = Stream::default(); + } +} From 0c28dcd2c1fc48513cb984435802963e107c0383 Mon Sep 17 00:00:00 2001 From: kerthcet Date: Mon, 20 Jul 2026 17:24:15 +0100 Subject: [PATCH 2/9] fix error return handling Signed-off-by: kerthcet --- crates/mlx/examples/hello.rs | 12 ++-- crates/mlx/src/array.rs | 117 ++++++++++++++++++++++++----------- crates/mlx/src/error.rs | 97 +++++++++++++++++++++++++++++ crates/mlx/src/lib.rs | 2 + 4 files changed, 187 insertions(+), 41 deletions(-) create mode 100644 crates/mlx/src/error.rs diff --git a/crates/mlx/examples/hello.rs b/crates/mlx/examples/hello.rs index de62b83..62c1005 100644 --- a/crates/mlx/examples/hello.rs +++ b/crates/mlx/examples/hello.rs @@ -7,7 +7,7 @@ use mlx::{Array, Stream}; -fn main() { +fn main() -> mlx::Result<()> { println!("MLX version: {}", mlx::version()); let a = Array::from_slice(&[1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]); @@ -18,15 +18,17 @@ fn main() { let b = Array::from_slice(&[6.0f32, 5.0, 4.0, 3.0, 2.0, 1.0], &[2, 3]); - // Operators use the default stream. + // Operators use the default stream and panic on error. println!("a + b:\n{:?}", &a + &b); // Or steer the default onto a specific device once, up front. Stream::cpu().set_as_default(); - let total = a.sum(false, &Stream::default()); + let total = a.sum(false, &Stream::default())?; println!("sum(a) on CPU: {}", total.item::()); - // Explicit stream control is still available per-op. - let product = a.multiply(&b, &Stream::gpu()); + // Explicit stream control is still available per-op; `?` propagates errors. + let product = a.multiply(&b, &Stream::gpu())?; println!("a * b (GPU):\n{product:?}"); + + Ok(()) } diff --git a/crates/mlx/src/array.rs b/crates/mlx/src/array.rs index 4b9eedd..c8e15ce 100644 --- a/crates/mlx/src/array.rs +++ b/crates/mlx/src/array.rs @@ -6,6 +6,7 @@ use std::fmt; use mlx_sys as sys; use crate::dtype::ArrayElement; +use crate::error::{self, Result}; use crate::stream::Stream; /// An N-dimensional MLX array. @@ -123,56 +124,56 @@ impl Array { } /// Elementwise addition: `self + other`. - pub fn add(&self, other: &Array, stream: &Stream) -> Array { + pub fn add(&self, other: &Array, stream: &Stream) -> Result { self.binary_op(other, stream, sys::mlx_add) } /// Elementwise subtraction: `self - other`. - pub fn subtract(&self, other: &Array, stream: &Stream) -> Array { + pub fn subtract(&self, other: &Array, stream: &Stream) -> Result { self.binary_op(other, stream, sys::mlx_subtract) } /// Elementwise multiplication: `self * other`. - pub fn multiply(&self, other: &Array, stream: &Stream) -> Array { + pub fn multiply(&self, other: &Array, stream: &Stream) -> Result { self.binary_op(other, stream, sys::mlx_multiply) } /// Elementwise division: `self / other`. - pub fn divide(&self, other: &Array, stream: &Stream) -> Array { + pub fn divide(&self, other: &Array, stream: &Stream) -> Result { self.binary_op(other, stream, sys::mlx_divide) } /// Elementwise square root. - pub fn sqrt(&self, stream: &Stream) -> Array { + pub fn sqrt(&self, stream: &Stream) -> Result { self.unary_op(stream, sys::mlx_sqrt) } /// Elementwise exponential. - pub fn exp(&self, stream: &Stream) -> Array { + pub fn exp(&self, stream: &Stream) -> Result { self.unary_op(stream, sys::mlx_exp) } /// Elementwise absolute value. - pub fn abs(&self, stream: &Stream) -> Array { + pub fn abs(&self, stream: &Stream) -> Result { self.unary_op(stream, sys::mlx_abs) } /// Elementwise negation. - pub fn negative(&self, stream: &Stream) -> Array { + pub fn negative(&self, stream: &Stream) -> Result { self.unary_op(stream, sys::mlx_negative) } /// Sum of all elements, returning a scalar array. /// /// With `keepdims == false` the result is 0-dimensional. - pub fn sum(&self, keepdims: bool, stream: &Stream) -> Array { + pub fn sum(&self, keepdims: bool, stream: &Stream) -> Result { self.reduce_op(keepdims, stream, sys::mlx_sum) } /// Mean of all elements, returning a scalar array. /// /// With `keepdims == false` the result is 0-dimensional. - pub fn mean(&self, keepdims: bool, stream: &Stream) -> Array { + pub fn mean(&self, keepdims: bool, stream: &Stream) -> Result { self.reduce_op(keepdims, stream, sys::mlx_mean) } @@ -187,13 +188,12 @@ impl Array { sys::mlx_array, sys::mlx_stream, ) -> i32, - ) -> Array { + ) -> Result { + error::install(); let mut out = unsafe { sys::mlx_array_new() }; // SAFETY: all handles are valid; `op` writes the result into `out`. - unsafe { - op(&mut out, self.handle, other.as_raw(), stream.as_raw()); - Self::from_raw(out) - } + let status = unsafe { op(&mut out, self.handle, other.as_raw(), stream.as_raw()) }; + Self::from_op(out, status) } /// Shared plumbing for `res = op(a, stream)` unary ops. @@ -201,13 +201,12 @@ impl Array { &self, stream: &Stream, op: unsafe extern "C" fn(*mut sys::mlx_array, sys::mlx_array, sys::mlx_stream) -> i32, - ) -> Array { + ) -> Result { + error::install(); let mut out = unsafe { sys::mlx_array_new() }; // SAFETY: handle/stream are valid; `op` writes the result into `out`. - unsafe { - op(&mut out, self.handle, stream.as_raw()); - Self::from_raw(out) - } + let status = unsafe { op(&mut out, self.handle, stream.as_raw()) }; + Self::from_op(out, status) } /// Shared plumbing for `res = op(a, keepdims, stream)` full reductions. @@ -216,12 +215,27 @@ impl Array { keepdims: bool, stream: &Stream, op: unsafe extern "C" fn(*mut sys::mlx_array, sys::mlx_array, bool, sys::mlx_stream) -> i32, - ) -> Array { + ) -> Result { + error::install(); let mut out = unsafe { sys::mlx_array_new() }; // SAFETY: handle/stream are valid; `op` writes the result into `out`. - unsafe { - op(&mut out, self.handle, keepdims, stream.as_raw()); - Self::from_raw(out) + let status = unsafe { op(&mut out, self.handle, keepdims, stream.as_raw()) }; + Self::from_op(out, status) + } + + /// Wraps an op's `out` handle and status code into a `Result`. + /// + /// On failure, frees the (unused) `out` handle and returns the captured + /// MLX error message. + fn from_op(out: sys::mlx_array, status: i32) -> Result { + match error::check(status) { + Ok(()) => Ok(unsafe { Self::from_raw(out) }), + Err(e) => { + // SAFETY: `out` was created by mlx and is owned here; free it so + // the failed op doesn't leak. + unsafe { sys::mlx_array_free(out) }; + Err(e) + } } } } @@ -248,8 +262,11 @@ impl Drop for Array { } // Arithmetic operators run on the current default stream (see -// [`Stream::set_as_default`]). For explicit stream control, call the inherent -// methods (`a.add(&b, &stream)`) instead. +// [`Stream::set_as_default`]). For explicit stream control — and to handle +// errors — call the inherent methods (`a.add(&b, &stream)?`) instead. +// +// Operators cannot return `Result`, so they **panic** if the underlying op +// fails (e.g. incompatible shapes). Use the methods when failure is possible. // // Implemented on `&Array` so operands are borrowed, not consumed: `&a + &b` // leaves both arrays usable afterwards. @@ -260,6 +277,7 @@ macro_rules! impl_binop { type Output = Array; fn $method(self, rhs: &Array) -> Array { self.$op(rhs, &Stream::default()) + .expect(concat!("Array::", stringify!($op), " failed")) } } )* @@ -277,6 +295,7 @@ impl std::ops::Neg for &Array { type Output = Array; fn neg(self) -> Array { self.negative(&Stream::default()) + .expect("Array::negative failed") } } @@ -329,21 +348,33 @@ mod tests { let a = Array::from_slice(&[1.0f32, 2.0, 3.0], &[3]); let b = Array::from_slice(&[4.0f32, 5.0, 6.0], &[3]); - assert_eq!(a.add(&b, &s).to_vec::(), vec![5.0, 7.0, 9.0]); - assert_eq!(b.subtract(&a, &s).to_vec::(), vec![3.0, 3.0, 3.0]); - assert_eq!(a.multiply(&b, &s).to_vec::(), vec![4.0, 10.0, 18.0]); - assert_eq!(b.divide(&a, &s).to_vec::(), vec![4.0, 2.5, 2.0]); + assert_eq!(a.add(&b, &s).unwrap().to_vec::(), vec![5.0, 7.0, 9.0]); + assert_eq!( + b.subtract(&a, &s).unwrap().to_vec::(), + vec![3.0, 3.0, 3.0] + ); + assert_eq!( + a.multiply(&b, &s).unwrap().to_vec::(), + vec![4.0, 10.0, 18.0] + ); + assert_eq!( + b.divide(&a, &s).unwrap().to_vec::(), + vec![4.0, 2.5, 2.0] + ); } #[test] fn unary_ops_compute_elementwise() { let s = Stream::cpu(); let a = Array::from_slice(&[1.0f32, 4.0, 9.0], &[3]); - assert_eq!(a.sqrt(&s).to_vec::(), vec![1.0, 2.0, 3.0]); + assert_eq!(a.sqrt(&s).unwrap().to_vec::(), vec![1.0, 2.0, 3.0]); let b = Array::from_slice(&[-1.0f32, 2.0, -3.0], &[3]); - assert_eq!(b.abs(&s).to_vec::(), vec![1.0, 2.0, 3.0]); - assert_eq!(b.negative(&s).to_vec::(), vec![1.0, -2.0, 3.0]); + assert_eq!(b.abs(&s).unwrap().to_vec::(), vec![1.0, 2.0, 3.0]); + assert_eq!( + b.negative(&s).unwrap().to_vec::(), + vec![1.0, -2.0, 3.0] + ); } #[test] @@ -351,22 +382,36 @@ mod tests { let s = Stream::cpu(); let a = Array::from_slice(&[1.0f32, 2.0, 3.0, 4.0], &[4]); - let sum = a.sum(false, &s); + let sum = a.sum(false, &s).unwrap(); assert_eq!(sum.ndim(), 0); assert_eq!(sum.item::(), 10.0); - assert_eq!(a.mean(false, &s).item::(), 2.5); + assert_eq!(a.mean(false, &s).unwrap().item::(), 2.5); } #[test] fn keepdims_retains_rank() { let s = Stream::cpu(); let a = Array::from_slice(&[1.0f32, 2.0, 3.0, 4.0], &[2, 2]); - let sum = a.sum(true, &s); + let sum = a.sum(true, &s).unwrap(); assert_eq!(sum.shape(), vec![1, 1]); assert_eq!(sum.item::(), 10.0); } + #[test] + fn incompatible_shapes_return_err() { + let s = Stream::cpu(); + let a = Array::from_slice(&[1.0f32, 2.0, 3.0], &[3]); + let b = Array::from_slice(&[1.0f32, 2.0], &[2]); + // Broadcasting [3] against [2] is invalid; MLX should report an error + // rather than aborting the process. + let err = a.add(&b, &s).unwrap_err(); + assert!( + !err.message().is_empty(), + "expected a non-empty error message" + ); + } + #[test] fn item_reads_scalar() { let a = Array::from_slice(&[42.0f32], &[]); diff --git a/crates/mlx/src/error.rs b/crates/mlx/src/error.rs new file mode 100644 index 0000000..88e0185 --- /dev/null +++ b/crates/mlx/src/error.rs @@ -0,0 +1,97 @@ +//! Error handling for MLX operations. +//! +//! MLX's C API reports failures two ways: an operation returns a non-zero +//! status code, and it routes a message to a globally-installed error handler. +//! The *default* handler prints to stderr and calls `exit(-1)` — fatal for a +//! library. On first use we install our own handler that captures the message +//! into thread-local storage instead, so failures surface as [`Error`] values. + +use std::cell::RefCell; +use std::ffi::{CStr, c_char, c_void}; +use std::fmt; +use std::ptr; +use std::sync::Once; + +use mlx_sys as sys; + +/// An error returned by an MLX operation. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Error { + message: String, +} + +impl Error { + pub(crate) fn new(message: impl Into) -> Self { + Self { + message: message.into(), + } + } + + /// The message MLX reported for this failure. + pub fn message(&self) -> &str { + &self.message + } +} + +impl fmt::Display for Error { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "MLX error: {}", self.message) + } +} + +impl std::error::Error for Error {} + +/// A `Result` whose error type is an MLX [`Error`]. +pub type Result = std::result::Result; + +thread_local! { + /// The most recent message from the MLX error handler, on this thread. + static LAST_ERROR: RefCell> = const { RefCell::new(None) }; +} + +/// The handler MLX invokes on failure. Called synchronously on the same thread +/// as the failing op, so thread-local capture is race-free. +unsafe extern "C" fn error_handler(msg: *const c_char, _data: *mut c_void) { + if msg.is_null() { + return; + } + // SAFETY: mlx passes a valid, NUL-terminated C string. + let message = unsafe { CStr::from_ptr(msg) } + .to_string_lossy() + .into_owned(); + LAST_ERROR.with(|e| *e.borrow_mut() = Some(message)); +} + +static INSTALL: Once = Once::new(); + +/// Installs our capturing error handler, replacing MLX's fatal default. +/// +/// Idempotent and cheap to call before every operation. +pub(crate) fn install() { + INSTALL.call_once(|| { + // SAFETY: `error_handler` has the required C signature; no user data or + // destructor is needed. + unsafe { + sys::mlx_set_error_handler(Some(error_handler), ptr::null_mut(), None); + } + }); +} + +/// Removes and returns the last captured error message on this thread. +pub(crate) fn take() -> Option { + LAST_ERROR.with(|e| e.borrow_mut().take()) +} + +/// Converts a status code into a `Result`, attaching the captured message. +/// +/// A non-zero `status` means the op failed; the message (if any) comes from the +/// handler that fired during the call. +pub(crate) fn check(status: i32) -> Result<()> { + if status == 0 { + Ok(()) + } else { + Err(Error::new( + take().unwrap_or_else(|| "unknown MLX error".to_string()), + )) + } +} diff --git a/crates/mlx/src/lib.rs b/crates/mlx/src/lib.rs index 1c4812b..a1710fa 100644 --- a/crates/mlx/src/lib.rs +++ b/crates/mlx/src/lib.rs @@ -5,10 +5,12 @@ mod array; mod dtype; +mod error; mod stream; pub use array::Array; pub use dtype::ArrayElement; +pub use error::{Error, Result}; pub use stream::Stream; /// Returns the version string of the underlying MLX library. From f39d4841fcf3044ad30c8665566933a3edf25f87 Mon Sep 17 00:00:00 2001 From: kerthcet Date: Mon, 20 Jul 2026 21:06:59 +0100 Subject: [PATCH 3/9] fix error Signed-off-by: kerthcet --- crates/mlx/src/array.rs | 25 ++++++++++++++++++++++--- crates/mlx/src/error.rs | 7 +++++++ crates/mlx/src/stream.rs | 4 ++++ 3 files changed, 33 insertions(+), 3 deletions(-) diff --git a/crates/mlx/src/array.rs b/crates/mlx/src/array.rs index c8e15ce..21e189d 100644 --- a/crates/mlx/src/array.rs +++ b/crates/mlx/src/array.rs @@ -45,6 +45,7 @@ impl Array { "data length {} does not match shape product {expected}", data.len() ); + error::install(); // SAFETY: pointers/len are valid for the duration of the call; mlx // copies the data into its own buffer. let handle = unsafe { @@ -82,6 +83,7 @@ impl Array { /// MLX is lazy: ops build a graph and only compute when the result is /// needed. `eval` materializes the values now. pub fn eval(&self) { + error::install(); // SAFETY: handle is valid for the lifetime of `self`. unsafe { sys::mlx_array_eval(self.handle); @@ -117,10 +119,18 @@ impl Array { "array dtype does not match requested element type" ); let len = self.size(); - // SAFETY: dtype matches `T` (checked above), so the pointer is valid for - // `size` contiguous `T` until the array is mutated or freed. + // `from_raw_parts` requires a non-null, aligned pointer even for a + // zero-length slice, but mlx may return null for an empty array. + if len == 0 { + return Vec::new(); + } + // SAFETY: dtype matches `T` (checked above), so mlx guarantees `len` + // contiguous, aligned `T` at `ptr`, valid until the array is mutated or + // freed. We only read (and copy out of) the slice within this call, so + // the borrow cannot outlive the buffer. `T: Copy`, so `to_vec` is a + // single bulk copy rather than `len` individual derefs. let ptr = unsafe { T::data_ptr(self.handle) }; - (0..len).map(|i| unsafe { *ptr.add(i) }).collect() + unsafe { std::slice::from_raw_parts(ptr, len) }.to_vec() } /// Elementwise addition: `self + other`. @@ -242,6 +252,7 @@ impl Array { impl fmt::Debug for Array { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + error::install(); // SAFETY: a freshly-created string handle is written by tostring. let mut s = unsafe { sys::mlx_string_new() }; unsafe { sys::mlx_array_tostring(&mut s, self.handle) }; @@ -435,6 +446,14 @@ mod tests { let _ = ints.to_vec::(); } + #[test] + fn to_vec_of_empty_array_is_empty() { + // A zero-element array must not deref a (possibly null) data pointer. + let empty = Array::from_slice::(&[], &[0]); + assert_eq!(empty.size(), 0); + assert!(empty.to_vec::().is_empty()); + } + #[test] fn operators_match_methods() { // Operators run on the default stream; results should equal the diff --git a/crates/mlx/src/error.rs b/crates/mlx/src/error.rs index 88e0185..2f833c3 100644 --- a/crates/mlx/src/error.rs +++ b/crates/mlx/src/error.rs @@ -86,7 +86,14 @@ pub(crate) fn take() -> Option { /// /// A non-zero `status` means the op failed; the message (if any) comes from the /// handler that fired during the call. +/// +/// Defensively ensures our handler is installed. This cannot rescue the op +/// whose `status` we are checking — that call has already returned — but it +/// guarantees any op reaching this central conversion point leaves the +/// process-wide handler in the non-fatal state, so MLX's default `exit(-1)` +/// handler can never linger even if a call site forgets to call [`install`]. pub(crate) fn check(status: i32) -> Result<()> { + install(); if status == 0 { Ok(()) } else { diff --git a/crates/mlx/src/stream.rs b/crates/mlx/src/stream.rs index df628cd..e09bb6c 100644 --- a/crates/mlx/src/stream.rs +++ b/crates/mlx/src/stream.rs @@ -14,6 +14,7 @@ pub struct Stream { impl Stream { /// The default stream on the GPU device. pub fn gpu() -> Self { + crate::error::install(); // SAFETY: constructor returns an owned handle. let handle = unsafe { sys::mlx_default_gpu_stream_new() }; Self { handle } @@ -21,6 +22,7 @@ impl Stream { /// The default stream on the CPU device. pub fn cpu() -> Self { + crate::error::install(); // SAFETY: constructor returns an owned handle. let handle = unsafe { sys::mlx_default_cpu_stream_new() }; Self { handle } @@ -39,6 +41,7 @@ impl Stream { /// current default. Call this once at startup to steer them onto a chosen /// device. pub fn set_as_default(&self) { + crate::error::install(); // SAFETY: `handle` is a valid stream owned by `self`; mlx copies out of // the pointers we pass and out of the handle. unsafe { @@ -58,6 +61,7 @@ impl Default for Stream { /// a GPU backend (Metal) is available, otherwise the CPU. This only reads /// that default; use [`Stream::set_as_default`] to change it. fn default() -> Self { + crate::error::install(); // SAFETY: out-params are valid; mlx writes owned handles into them. let handle = unsafe { let mut dev = sys::mlx_device_new(); From 81d9f831351466a798ffb28c605eb84448d75e6e Mon Sep 17 00:00:00 2001 From: kerthcet Date: Mon, 20 Jul 2026 21:14:41 +0100 Subject: [PATCH 4/9] fix stale error Signed-off-by: kerthcet --- crates/mlx/src/error.rs | 3 +++ 1 file changed, 3 insertions(+) diff --git a/crates/mlx/src/error.rs b/crates/mlx/src/error.rs index 2f833c3..6a7575e 100644 --- a/crates/mlx/src/error.rs +++ b/crates/mlx/src/error.rs @@ -95,6 +95,9 @@ pub(crate) fn take() -> Option { pub(crate) fn check(status: i32) -> Result<()> { install(); if status == 0 { + // Clear any message left by an earlier call so it can never be + // misattributed to a later failure whose handler did not fire. + take(); Ok(()) } else { Err(Error::new( From 1e1f30409a91446a4aa5fab468dff9bb76dbd7d1 Mon Sep 17 00:00:00 2001 From: kerthcet Date: Mon, 20 Jul 2026 23:33:44 +0100 Subject: [PATCH 5/9] rename to mlxr Signed-off-by: kerthcet --- Cargo.lock | 6 +++--- Cargo.toml | 6 +++--- README.md | 12 ++++++------ crates/{mlx-sys => mlxr-sys}/Cargo.toml | 4 ++-- crates/{mlx-sys => mlxr-sys}/build.rs | 10 +++++----- crates/{mlx-sys => mlxr-sys}/src/lib.rs | 2 +- crates/{mlx => mlxr}/Cargo.toml | 8 ++++---- crates/{mlx => mlxr}/examples/hello.rs | 6 +++--- crates/{mlx => mlxr}/src/array.rs | 2 +- crates/{mlx => mlxr}/src/dtype.rs | 2 +- crates/{mlx => mlxr}/src/error.rs | 2 +- crates/{mlx => mlxr}/src/lib.rs | 10 +++++----- crates/{mlx => mlxr}/src/stream.rs | 2 +- 13 files changed, 36 insertions(+), 36 deletions(-) rename crates/{mlx-sys => mlxr-sys}/Cargo.toml (85%) rename crates/{mlx-sys => mlxr-sys}/build.rs (92%) rename crates/{mlx-sys => mlxr-sys}/src/lib.rs (90%) rename crates/{mlx => mlxr}/Cargo.toml (72%) rename crates/{mlx => mlxr}/examples/hello.rs (89%) rename crates/{mlx => mlxr}/src/array.rs (99%) rename crates/{mlx => mlxr}/src/dtype.rs (99%) rename crates/{mlx => mlxr}/src/error.rs (99%) rename crates/{mlx => mlxr}/src/lib.rs (73%) rename crates/{mlx => mlxr}/src/stream.rs (99%) diff --git a/Cargo.lock b/Cargo.lock index 9f8661d..7d1bdd6 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -144,14 +144,14 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a" [[package]] -name = "mlx" +name = "mlxr" version = "0.0.0" dependencies = [ - "mlx-sys", + "mlxr-sys", ] [[package]] -name = "mlx-sys" +name = "mlxr-sys" version = "0.0.0" dependencies = [ "bindgen", diff --git a/Cargo.toml b/Cargo.toml index 8333f53..312746f 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,12 +1,12 @@ [workspace] resolver = "2" -members = ["crates/mlx-sys", "crates/mlx"] +members = ["crates/mlxr-sys", "crates/mlxr"] [workspace.package] edition = "2024" license = "MIT" -repository = "https://github.com/InftyAI/mlx" +repository = "https://github.com/InftyAI/mlxr" authors = ["kerthcet"] [workspace.dependencies] -mlx-sys = { path = "crates/mlx-sys", version = "0.0.0" } +mlxr-sys = { path = "crates/mlxr-sys", version = "0.0.0" } diff --git a/README.md b/README.md index 359f17f..3e91cc8 100644 --- a/README.md +++ b/README.md @@ -1,4 +1,4 @@ -# MLX +# MLXR (MLX Rust) Rust bindings for Apple's [MLX](https://github.com/ml-explore/mlx) framework. @@ -9,10 +9,10 @@ Rust bindings for Apple's [MLX](https://github.com/ml-explore/mlx) framework. The bindings are layered, the same way as [MLX Swift](https://github.com/ml-explore/mlx-swift): ``` -mlx (safe, idiomatic Rust API) <- crates/mlx -mlx-sys (raw unsafe FFI bindings) <- crates/mlx-sys, generated by bindgen -mlx-c (Apple's C API for MLX) <- third_party/mlx-c (git submodule) -mlx (Apple's C++ framework) <- fetched by mlx-c's CMake build +mlxr (safe, idiomatic Rust API) <- crates/mlxr +mlxr-sys (raw unsafe FFI bindings) <- crates/mlxr-sys, generated by bindgen +mlx-c (Apple's C API for MLX) <- third_party/mlx-c (git submodule) +mlx (Apple's C++ framework) <- fetched by mlx-c's CMake build ``` MLX itself is C++, which Rust cannot bind to directly. We bind against @@ -32,7 +32,7 @@ The first build compiles MLX from source and takes several minutes. ## Examples -Runnable examples live in `crates/mlx/examples`: +Runnable examples live in `crates/mlxr/examples`: ```sh cargo run --example hello diff --git a/crates/mlx-sys/Cargo.toml b/crates/mlxr-sys/Cargo.toml similarity index 85% rename from crates/mlx-sys/Cargo.toml rename to crates/mlxr-sys/Cargo.toml index 40704c8..4d5e8d7 100644 --- a/crates/mlx-sys/Cargo.toml +++ b/crates/mlxr-sys/Cargo.toml @@ -1,8 +1,8 @@ [package] -name = "mlx-sys" +name = "mlxr-sys" version = "0.0.0" # Edition 2021: bindgen 0.70 emits plain `extern "C"` blocks, which edition -# 2024 rejects (it requires `unsafe extern`). The safe `mlx` crate stays 2024. +# 2024 rejects (it requires `unsafe extern`). The safe `mlxr` crate stays 2024. edition = "2021" license.workspace = true repository.workspace = true diff --git a/crates/mlx-sys/build.rs b/crates/mlxr-sys/build.rs similarity index 92% rename from crates/mlx-sys/build.rs rename to crates/mlxr-sys/build.rs index 65d9315..f83dc31 100644 --- a/crates/mlx-sys/build.rs +++ b/crates/mlxr-sys/build.rs @@ -1,4 +1,4 @@ -//! Build script for `mlx-sys`. +//! Build script for `mlxr-sys`. //! //! 1. Builds the vendored `mlx-c` C API with CMake. `mlx-c` uses CMake //! `FetchContent` to download and build MLX itself, so both `libmlxc` and @@ -13,11 +13,11 @@ use std::path::PathBuf; fn main() { if !cfg!(target_os = "macos") { - panic!("mlx-sys currently only supports macOS on Apple Silicon"); + panic!("mlxr-sys currently only supports macOS on Apple Silicon"); } let manifest_dir = PathBuf::from(env::var("CARGO_MANIFEST_DIR").unwrap()); - // Workspace layout: /crates/mlx-sys -> /third_party/mlx-c + // Workspace layout: /crates/mlxr-sys -> /third_party/mlx-c let mlx_c_dir = manifest_dir .join("../../third_party/mlx-c") .canonicalize() @@ -39,12 +39,12 @@ fn main() { } // Build into a fixed dir under `target//` rather than the default - // per-fingerprint `OUT_DIR`. `cargo clippy` fingerprints `mlx-sys` + // per-fingerprint `OUT_DIR`. `cargo clippy` fingerprints `mlxr-sys` // differently from `cargo build`/`cargo test`, so an `OUT_DIR`-based build // would recompile MLX from scratch for each (~2.5 min of C++). A shared dir // lets them all reuse the same CMake build. // - // OUT_DIR = //build/mlx-sys-/out; three parents up is + // OUT_DIR = //build/mlxr-sys-/out; three parents up is // /. let out_dir = PathBuf::from(env::var("OUT_DIR").unwrap()); let profile_dir = out_dir diff --git a/crates/mlx-sys/src/lib.rs b/crates/mlxr-sys/src/lib.rs similarity index 90% rename from crates/mlx-sys/src/lib.rs rename to crates/mlxr-sys/src/lib.rs index 5269012..1f25fba 100644 --- a/crates/mlx-sys/src/lib.rs +++ b/crates/mlxr-sys/src/lib.rs @@ -3,7 +3,7 @@ //! //! These bindings are generated at build time by `bindgen` from the pinned //! `third_party/mlx-c` submodule. This crate is not meant to be used directly; -//! prefer the safe `mlx` crate that wraps it. +//! prefer the safe `mlxr` crate that wraps it. #![allow(non_upper_case_globals)] #![allow(non_camel_case_types)] #![allow(non_snake_case)] diff --git a/crates/mlx/Cargo.toml b/crates/mlxr/Cargo.toml similarity index 72% rename from crates/mlx/Cargo.toml rename to crates/mlxr/Cargo.toml index bebe917..83a4ecf 100644 --- a/crates/mlx/Cargo.toml +++ b/crates/mlxr/Cargo.toml @@ -1,5 +1,5 @@ [package] -name = "mlx" +name = "mlxr" version = "0.0.0" edition.workspace = true license.workspace = true @@ -9,8 +9,8 @@ description = "Safe, idiomatic Rust bindings for Apple's MLX array framework." [features] default = ["metal", "accelerate"] -metal = ["mlx-sys/metal"] -accelerate = ["mlx-sys/accelerate"] +metal = ["mlxr-sys/metal"] +accelerate = ["mlxr-sys/accelerate"] [dependencies] -mlx-sys.workspace = true +mlxr-sys.workspace = true diff --git a/crates/mlx/examples/hello.rs b/crates/mlxr/examples/hello.rs similarity index 89% rename from crates/mlx/examples/hello.rs rename to crates/mlxr/examples/hello.rs index 62c1005..1769a7b 100644 --- a/crates/mlx/examples/hello.rs +++ b/crates/mlxr/examples/hello.rs @@ -5,10 +5,10 @@ //! cargo run --example hello //! ``` -use mlx::{Array, Stream}; +use mlxr::{Array, Stream}; -fn main() -> mlx::Result<()> { - println!("MLX version: {}", mlx::version()); +fn main() -> mlxr::Result<()> { + println!("MLX version: {}", mlxr::version()); let a = Array::from_slice(&[1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]); println!("shape: {:?}", a.shape()); diff --git a/crates/mlx/src/array.rs b/crates/mlxr/src/array.rs similarity index 99% rename from crates/mlx/src/array.rs rename to crates/mlxr/src/array.rs index 21e189d..d37841f 100644 --- a/crates/mlx/src/array.rs +++ b/crates/mlxr/src/array.rs @@ -3,7 +3,7 @@ use std::ffi::CStr; use std::fmt; -use mlx_sys as sys; +use mlxr_sys as sys; use crate::dtype::ArrayElement; use crate::error::{self, Result}; diff --git a/crates/mlx/src/dtype.rs b/crates/mlxr/src/dtype.rs similarity index 99% rename from crates/mlx/src/dtype.rs rename to crates/mlxr/src/dtype.rs index 63b4776..25c9fbb 100644 --- a/crates/mlx/src/dtype.rs +++ b/crates/mlxr/src/dtype.rs @@ -1,6 +1,6 @@ //! Mapping between Rust primitive types and MLX data types. -use mlx_sys as sys; +use mlxr_sys as sys; mod sealed { pub trait Sealed {} diff --git a/crates/mlx/src/error.rs b/crates/mlxr/src/error.rs similarity index 99% rename from crates/mlx/src/error.rs rename to crates/mlxr/src/error.rs index 6a7575e..cb375ee 100644 --- a/crates/mlx/src/error.rs +++ b/crates/mlxr/src/error.rs @@ -12,7 +12,7 @@ use std::fmt; use std::ptr; use std::sync::Once; -use mlx_sys as sys; +use mlxr_sys as sys; /// An error returned by an MLX operation. #[derive(Debug, Clone, PartialEq, Eq)] diff --git a/crates/mlx/src/lib.rs b/crates/mlxr/src/lib.rs similarity index 73% rename from crates/mlx/src/lib.rs rename to crates/mlxr/src/lib.rs index a1710fa..b5117a2 100644 --- a/crates/mlx/src/lib.rs +++ b/crates/mlxr/src/lib.rs @@ -1,5 +1,5 @@ //! Safe, idiomatic Rust bindings for Apple's [MLX](https://github.com/ml-explore/mlx) -//! array framework, built on top of the [`mlx-sys`] FFI layer. +//! array framework, built on top of the [`mlxr-sys`] FFI layer. //! //! This crate is Apple Silicon (macOS) only. @@ -18,12 +18,12 @@ pub fn version() -> String { use std::ffi::CStr; // SAFETY: standard mlx-c string-handle dance; all handles are freed. unsafe { - let mut s = mlx_sys::mlx_string_new(); - mlx_sys::mlx_version(&mut s); - let v = CStr::from_ptr(mlx_sys::mlx_string_data(s)) + let mut s = mlxr_sys::mlx_string_new(); + mlxr_sys::mlx_version(&mut s); + let v = CStr::from_ptr(mlxr_sys::mlx_string_data(s)) .to_string_lossy() .into_owned(); - mlx_sys::mlx_string_free(s); + mlxr_sys::mlx_string_free(s); v } } diff --git a/crates/mlx/src/stream.rs b/crates/mlxr/src/stream.rs similarity index 99% rename from crates/mlx/src/stream.rs rename to crates/mlxr/src/stream.rs index e09bb6c..52cb42a 100644 --- a/crates/mlx/src/stream.rs +++ b/crates/mlxr/src/stream.rs @@ -4,7 +4,7 @@ //! (CPU or GPU). Most ops take a stream argument. MLX's own default is the GPU //! when a GPU backend is available (as on Apple Silicon), else the CPU. -use mlx_sys as sys; +use mlxr_sys as sys; /// An execution stream bound to a device. pub struct Stream { From 5cff1b1622031cdbc0ef64839b2bbc5bcafd89c6 Mon Sep 17 00:00:00 2001 From: kerthcet Date: Mon, 20 Jul 2026 23:36:35 +0100 Subject: [PATCH 6/9] log error Signed-off-by: kerthcet --- crates/mlxr/src/array.rs | 17 ++++++++++++++--- 1 file changed, 14 insertions(+), 3 deletions(-) diff --git a/crates/mlxr/src/array.rs b/crates/mlxr/src/array.rs index d37841f..de90488 100644 --- a/crates/mlxr/src/array.rs +++ b/crates/mlxr/src/array.rs @@ -287,8 +287,9 @@ macro_rules! impl_binop { impl std::ops::$trait for &Array { type Output = Array; fn $method(self, rhs: &Array) -> Array { - self.$op(rhs, &Stream::default()) - .expect(concat!("Array::", stringify!($op), " failed")) + self.$op(rhs, &Stream::default()).unwrap_or_else(|e| { + panic!(concat!("Array::", stringify!($op), " failed: {}"), e) + }) } } )* @@ -306,7 +307,7 @@ impl std::ops::Neg for &Array { type Output = Array; fn neg(self) -> Array { self.negative(&Stream::default()) - .expect("Array::negative failed") + .unwrap_or_else(|e| panic!("Array::negative failed: {e}")) } } @@ -470,4 +471,14 @@ mod tests { // Operands are borrowed, so `a` is still usable here. assert_eq!(a.to_vec::(), vec![10.0, 20.0, 30.0]); } + + #[test] + #[should_panic(expected = "Array::add failed: MLX error:")] + fn operator_panic_carries_mlx_message() { + // A failing operator panics with both the op name and the underlying + // MLX diagnostic, not a generic message. + let a = Array::from_slice(&[1.0f32, 2.0, 3.0], &[3]); + let b = Array::from_slice(&[1.0f32, 2.0], &[2]); + let _ = &a + &b; + } } From 39f1074ae6b103bd189b634ffd7915bea35522dd Mon Sep 17 00:00:00 2001 From: kerthcet Date: Mon, 20 Jul 2026 23:40:46 +0100 Subject: [PATCH 7/9] revert the folder name Signed-off-by: kerthcet --- Cargo.toml | 4 ++-- README.md | 6 +++--- crates/{mlxr-sys => mlx-sys}/Cargo.toml | 0 crates/{mlxr-sys => mlx-sys}/build.rs | 2 +- crates/{mlxr-sys => mlx-sys}/src/lib.rs | 0 crates/{mlxr => mlx}/Cargo.toml | 0 crates/{mlxr => mlx}/examples/hello.rs | 0 crates/{mlxr => mlx}/src/array.rs | 0 crates/{mlxr => mlx}/src/dtype.rs | 0 crates/{mlxr => mlx}/src/error.rs | 0 crates/{mlxr => mlx}/src/lib.rs | 0 crates/{mlxr => mlx}/src/stream.rs | 0 12 files changed, 6 insertions(+), 6 deletions(-) rename crates/{mlxr-sys => mlx-sys}/Cargo.toml (100%) rename crates/{mlxr-sys => mlx-sys}/build.rs (97%) rename crates/{mlxr-sys => mlx-sys}/src/lib.rs (100%) rename crates/{mlxr => mlx}/Cargo.toml (100%) rename crates/{mlxr => mlx}/examples/hello.rs (100%) rename crates/{mlxr => mlx}/src/array.rs (100%) rename crates/{mlxr => mlx}/src/dtype.rs (100%) rename crates/{mlxr => mlx}/src/error.rs (100%) rename crates/{mlxr => mlx}/src/lib.rs (100%) rename crates/{mlxr => mlx}/src/stream.rs (100%) diff --git a/Cargo.toml b/Cargo.toml index 312746f..93914cb 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [workspace] resolver = "2" -members = ["crates/mlxr-sys", "crates/mlxr"] +members = ["crates/mlx-sys", "crates/mlx"] [workspace.package] edition = "2024" @@ -9,4 +9,4 @@ repository = "https://github.com/InftyAI/mlxr" authors = ["kerthcet"] [workspace.dependencies] -mlxr-sys = { path = "crates/mlxr-sys", version = "0.0.0" } +mlxr-sys = { path = "crates/mlx-sys", version = "0.0.0" } diff --git a/README.md b/README.md index 3e91cc8..68b86c9 100644 --- a/README.md +++ b/README.md @@ -9,8 +9,8 @@ Rust bindings for Apple's [MLX](https://github.com/ml-explore/mlx) framework. The bindings are layered, the same way as [MLX Swift](https://github.com/ml-explore/mlx-swift): ``` -mlxr (safe, idiomatic Rust API) <- crates/mlxr -mlxr-sys (raw unsafe FFI bindings) <- crates/mlxr-sys, generated by bindgen +mlxr (safe, idiomatic Rust API) <- crates/mlx +mlxr-sys (raw unsafe FFI bindings) <- crates/mlx-sys, generated by bindgen mlx-c (Apple's C API for MLX) <- third_party/mlx-c (git submodule) mlx (Apple's C++ framework) <- fetched by mlx-c's CMake build ``` @@ -32,7 +32,7 @@ The first build compiles MLX from source and takes several minutes. ## Examples -Runnable examples live in `crates/mlxr/examples`: +Runnable examples live in `crates/mlx/examples`: ```sh cargo run --example hello diff --git a/crates/mlxr-sys/Cargo.toml b/crates/mlx-sys/Cargo.toml similarity index 100% rename from crates/mlxr-sys/Cargo.toml rename to crates/mlx-sys/Cargo.toml diff --git a/crates/mlxr-sys/build.rs b/crates/mlx-sys/build.rs similarity index 97% rename from crates/mlxr-sys/build.rs rename to crates/mlx-sys/build.rs index f83dc31..2332a7a 100644 --- a/crates/mlxr-sys/build.rs +++ b/crates/mlx-sys/build.rs @@ -17,7 +17,7 @@ fn main() { } let manifest_dir = PathBuf::from(env::var("CARGO_MANIFEST_DIR").unwrap()); - // Workspace layout: /crates/mlxr-sys -> /third_party/mlx-c + // Workspace layout: /crates/mlx-sys -> /third_party/mlx-c let mlx_c_dir = manifest_dir .join("../../third_party/mlx-c") .canonicalize() diff --git a/crates/mlxr-sys/src/lib.rs b/crates/mlx-sys/src/lib.rs similarity index 100% rename from crates/mlxr-sys/src/lib.rs rename to crates/mlx-sys/src/lib.rs diff --git a/crates/mlxr/Cargo.toml b/crates/mlx/Cargo.toml similarity index 100% rename from crates/mlxr/Cargo.toml rename to crates/mlx/Cargo.toml diff --git a/crates/mlxr/examples/hello.rs b/crates/mlx/examples/hello.rs similarity index 100% rename from crates/mlxr/examples/hello.rs rename to crates/mlx/examples/hello.rs diff --git a/crates/mlxr/src/array.rs b/crates/mlx/src/array.rs similarity index 100% rename from crates/mlxr/src/array.rs rename to crates/mlx/src/array.rs diff --git a/crates/mlxr/src/dtype.rs b/crates/mlx/src/dtype.rs similarity index 100% rename from crates/mlxr/src/dtype.rs rename to crates/mlx/src/dtype.rs diff --git a/crates/mlxr/src/error.rs b/crates/mlx/src/error.rs similarity index 100% rename from crates/mlxr/src/error.rs rename to crates/mlx/src/error.rs diff --git a/crates/mlxr/src/lib.rs b/crates/mlx/src/lib.rs similarity index 100% rename from crates/mlxr/src/lib.rs rename to crates/mlx/src/lib.rs diff --git a/crates/mlxr/src/stream.rs b/crates/mlx/src/stream.rs similarity index 100% rename from crates/mlxr/src/stream.rs rename to crates/mlx/src/stream.rs From 267f18c282d8c5aae14c84eef2f6643ad7e180aa Mon Sep 17 00:00:00 2001 From: kerthcet Date: Tue, 21 Jul 2026 00:03:40 +0100 Subject: [PATCH 8/9] rename to mlxcore Signed-off-by: kerthcet --- Cargo.lock | 6 +++--- Cargo.toml | 6 +++--- README.md | 8 ++++---- crates/{mlx-sys => mlxcore-sys}/Cargo.toml | 4 ++-- crates/{mlx-sys => mlxcore-sys}/build.rs | 10 +++++----- crates/{mlx-sys => mlxcore-sys}/src/lib.rs | 2 +- crates/{mlx => mlxcore}/Cargo.toml | 8 ++++---- crates/{mlx => mlxcore}/examples/hello.rs | 6 +++--- crates/{mlx => mlxcore}/src/array.rs | 2 +- crates/{mlx => mlxcore}/src/dtype.rs | 2 +- crates/{mlx => mlxcore}/src/error.rs | 2 +- crates/{mlx => mlxcore}/src/lib.rs | 10 +++++----- crates/{mlx => mlxcore}/src/stream.rs | 2 +- 13 files changed, 34 insertions(+), 34 deletions(-) rename crates/{mlx-sys => mlxcore-sys}/Cargo.toml (84%) rename crates/{mlx-sys => mlxcore-sys}/build.rs (92%) rename crates/{mlx-sys => mlxcore-sys}/src/lib.rs (89%) rename crates/{mlx => mlxcore}/Cargo.toml (70%) rename crates/{mlx => mlxcore}/examples/hello.rs (88%) rename crates/{mlx => mlxcore}/src/array.rs (99%) rename crates/{mlx => mlxcore}/src/dtype.rs (99%) rename crates/{mlx => mlxcore}/src/error.rs (99%) rename crates/{mlx => mlxcore}/src/lib.rs (72%) rename crates/{mlx => mlxcore}/src/stream.rs (99%) diff --git a/Cargo.lock b/Cargo.lock index 7d1bdd6..dcb36b8 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -144,14 +144,14 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a" [[package]] -name = "mlxr" +name = "mlxcore" version = "0.0.0" dependencies = [ - "mlxr-sys", + "mlxcore-sys", ] [[package]] -name = "mlxr-sys" +name = "mlxcore-sys" version = "0.0.0" dependencies = [ "bindgen", diff --git a/Cargo.toml b/Cargo.toml index 93914cb..32597c0 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,12 +1,12 @@ [workspace] resolver = "2" -members = ["crates/mlx-sys", "crates/mlx"] +members = ["crates/mlxcore-sys", "crates/mlxcore"] [workspace.package] edition = "2024" license = "MIT" -repository = "https://github.com/InftyAI/mlxr" +repository = "https://github.com/InftyAI/mlx" authors = ["kerthcet"] [workspace.dependencies] -mlxr-sys = { path = "crates/mlx-sys", version = "0.0.0" } +mlxcore-sys = { path = "crates/mlxcore-sys", version = "0.0.0" } diff --git a/README.md b/README.md index 68b86c9..d063479 100644 --- a/README.md +++ b/README.md @@ -1,4 +1,4 @@ -# MLXR (MLX Rust) +# MLX Rust bindings for Apple's [MLX](https://github.com/ml-explore/mlx) framework. @@ -9,8 +9,8 @@ Rust bindings for Apple's [MLX](https://github.com/ml-explore/mlx) framework. The bindings are layered, the same way as [MLX Swift](https://github.com/ml-explore/mlx-swift): ``` -mlxr (safe, idiomatic Rust API) <- crates/mlx -mlxr-sys (raw unsafe FFI bindings) <- crates/mlx-sys, generated by bindgen +mlxcore (safe, idiomatic Rust API) <- crates/mlxcore +mlxcore-sys (raw unsafe FFI bindings) <- crates/mlxcore-sys, generated by bindgen mlx-c (Apple's C API for MLX) <- third_party/mlx-c (git submodule) mlx (Apple's C++ framework) <- fetched by mlx-c's CMake build ``` @@ -32,7 +32,7 @@ The first build compiles MLX from source and takes several minutes. ## Examples -Runnable examples live in `crates/mlx/examples`: +Runnable examples live in `crates/mlxcore/examples`: ```sh cargo run --example hello diff --git a/crates/mlx-sys/Cargo.toml b/crates/mlxcore-sys/Cargo.toml similarity index 84% rename from crates/mlx-sys/Cargo.toml rename to crates/mlxcore-sys/Cargo.toml index 4d5e8d7..1ad9adc 100644 --- a/crates/mlx-sys/Cargo.toml +++ b/crates/mlxcore-sys/Cargo.toml @@ -1,8 +1,8 @@ [package] -name = "mlxr-sys" +name = "mlxcore-sys" version = "0.0.0" # Edition 2021: bindgen 0.70 emits plain `extern "C"` blocks, which edition -# 2024 rejects (it requires `unsafe extern`). The safe `mlxr` crate stays 2024. +# 2024 rejects (it requires `unsafe extern`). The safe `mlxcore` crate stays 2024. edition = "2021" license.workspace = true repository.workspace = true diff --git a/crates/mlx-sys/build.rs b/crates/mlxcore-sys/build.rs similarity index 92% rename from crates/mlx-sys/build.rs rename to crates/mlxcore-sys/build.rs index 2332a7a..40e310c 100644 --- a/crates/mlx-sys/build.rs +++ b/crates/mlxcore-sys/build.rs @@ -1,4 +1,4 @@ -//! Build script for `mlxr-sys`. +//! Build script for `mlxcore-sys`. //! //! 1. Builds the vendored `mlx-c` C API with CMake. `mlx-c` uses CMake //! `FetchContent` to download and build MLX itself, so both `libmlxc` and @@ -13,11 +13,11 @@ use std::path::PathBuf; fn main() { if !cfg!(target_os = "macos") { - panic!("mlxr-sys currently only supports macOS on Apple Silicon"); + panic!("mlxcore-sys currently only supports macOS on Apple Silicon"); } let manifest_dir = PathBuf::from(env::var("CARGO_MANIFEST_DIR").unwrap()); - // Workspace layout: /crates/mlx-sys -> /third_party/mlx-c + // Workspace layout: /crates/mlxcore-sys -> /third_party/mlx-c let mlx_c_dir = manifest_dir .join("../../third_party/mlx-c") .canonicalize() @@ -39,12 +39,12 @@ fn main() { } // Build into a fixed dir under `target//` rather than the default - // per-fingerprint `OUT_DIR`. `cargo clippy` fingerprints `mlxr-sys` + // per-fingerprint `OUT_DIR`. `cargo clippy` fingerprints `mlxcore-sys` // differently from `cargo build`/`cargo test`, so an `OUT_DIR`-based build // would recompile MLX from scratch for each (~2.5 min of C++). A shared dir // lets them all reuse the same CMake build. // - // OUT_DIR = //build/mlxr-sys-/out; three parents up is + // OUT_DIR = //build/mlxcore-sys-/out; three parents up is // /. let out_dir = PathBuf::from(env::var("OUT_DIR").unwrap()); let profile_dir = out_dir diff --git a/crates/mlx-sys/src/lib.rs b/crates/mlxcore-sys/src/lib.rs similarity index 89% rename from crates/mlx-sys/src/lib.rs rename to crates/mlxcore-sys/src/lib.rs index 1f25fba..2e4fcfc 100644 --- a/crates/mlx-sys/src/lib.rs +++ b/crates/mlxcore-sys/src/lib.rs @@ -3,7 +3,7 @@ //! //! These bindings are generated at build time by `bindgen` from the pinned //! `third_party/mlx-c` submodule. This crate is not meant to be used directly; -//! prefer the safe `mlxr` crate that wraps it. +//! prefer the safe `mlxcore` crate that wraps it. #![allow(non_upper_case_globals)] #![allow(non_camel_case_types)] #![allow(non_snake_case)] diff --git a/crates/mlx/Cargo.toml b/crates/mlxcore/Cargo.toml similarity index 70% rename from crates/mlx/Cargo.toml rename to crates/mlxcore/Cargo.toml index 83a4ecf..58f7a36 100644 --- a/crates/mlx/Cargo.toml +++ b/crates/mlxcore/Cargo.toml @@ -1,5 +1,5 @@ [package] -name = "mlxr" +name = "mlxcore" version = "0.0.0" edition.workspace = true license.workspace = true @@ -9,8 +9,8 @@ description = "Safe, idiomatic Rust bindings for Apple's MLX array framework." [features] default = ["metal", "accelerate"] -metal = ["mlxr-sys/metal"] -accelerate = ["mlxr-sys/accelerate"] +metal = ["mlxcore-sys/metal"] +accelerate = ["mlxcore-sys/accelerate"] [dependencies] -mlxr-sys.workspace = true +mlxcore-sys.workspace = true diff --git a/crates/mlx/examples/hello.rs b/crates/mlxcore/examples/hello.rs similarity index 88% rename from crates/mlx/examples/hello.rs rename to crates/mlxcore/examples/hello.rs index 1769a7b..74b0ede 100644 --- a/crates/mlx/examples/hello.rs +++ b/crates/mlxcore/examples/hello.rs @@ -5,10 +5,10 @@ //! cargo run --example hello //! ``` -use mlxr::{Array, Stream}; +use mlxcore::{Array, Stream}; -fn main() -> mlxr::Result<()> { - println!("MLX version: {}", mlxr::version()); +fn main() -> mlxcore::Result<()> { + println!("MLX version: {}", mlxcore::version()); let a = Array::from_slice(&[1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]); println!("shape: {:?}", a.shape()); diff --git a/crates/mlx/src/array.rs b/crates/mlxcore/src/array.rs similarity index 99% rename from crates/mlx/src/array.rs rename to crates/mlxcore/src/array.rs index de90488..998dbf8 100644 --- a/crates/mlx/src/array.rs +++ b/crates/mlxcore/src/array.rs @@ -3,7 +3,7 @@ use std::ffi::CStr; use std::fmt; -use mlxr_sys as sys; +use mlxcore_sys as sys; use crate::dtype::ArrayElement; use crate::error::{self, Result}; diff --git a/crates/mlx/src/dtype.rs b/crates/mlxcore/src/dtype.rs similarity index 99% rename from crates/mlx/src/dtype.rs rename to crates/mlxcore/src/dtype.rs index 25c9fbb..55851e3 100644 --- a/crates/mlx/src/dtype.rs +++ b/crates/mlxcore/src/dtype.rs @@ -1,6 +1,6 @@ //! Mapping between Rust primitive types and MLX data types. -use mlxr_sys as sys; +use mlxcore_sys as sys; mod sealed { pub trait Sealed {} diff --git a/crates/mlx/src/error.rs b/crates/mlxcore/src/error.rs similarity index 99% rename from crates/mlx/src/error.rs rename to crates/mlxcore/src/error.rs index cb375ee..2dc641a 100644 --- a/crates/mlx/src/error.rs +++ b/crates/mlxcore/src/error.rs @@ -12,7 +12,7 @@ use std::fmt; use std::ptr; use std::sync::Once; -use mlxr_sys as sys; +use mlxcore_sys as sys; /// An error returned by an MLX operation. #[derive(Debug, Clone, PartialEq, Eq)] diff --git a/crates/mlx/src/lib.rs b/crates/mlxcore/src/lib.rs similarity index 72% rename from crates/mlx/src/lib.rs rename to crates/mlxcore/src/lib.rs index b5117a2..ad8c59b 100644 --- a/crates/mlx/src/lib.rs +++ b/crates/mlxcore/src/lib.rs @@ -1,5 +1,5 @@ //! Safe, idiomatic Rust bindings for Apple's [MLX](https://github.com/ml-explore/mlx) -//! array framework, built on top of the [`mlxr-sys`] FFI layer. +//! array framework, built on top of the [`mlxcore-sys`] FFI layer. //! //! This crate is Apple Silicon (macOS) only. @@ -18,12 +18,12 @@ pub fn version() -> String { use std::ffi::CStr; // SAFETY: standard mlx-c string-handle dance; all handles are freed. unsafe { - let mut s = mlxr_sys::mlx_string_new(); - mlxr_sys::mlx_version(&mut s); - let v = CStr::from_ptr(mlxr_sys::mlx_string_data(s)) + let mut s = mlxcore_sys::mlx_string_new(); + mlxcore_sys::mlx_version(&mut s); + let v = CStr::from_ptr(mlxcore_sys::mlx_string_data(s)) .to_string_lossy() .into_owned(); - mlxr_sys::mlx_string_free(s); + mlxcore_sys::mlx_string_free(s); v } } diff --git a/crates/mlx/src/stream.rs b/crates/mlxcore/src/stream.rs similarity index 99% rename from crates/mlx/src/stream.rs rename to crates/mlxcore/src/stream.rs index 52cb42a..90d87f4 100644 --- a/crates/mlx/src/stream.rs +++ b/crates/mlxcore/src/stream.rs @@ -4,7 +4,7 @@ //! (CPU or GPU). Most ops take a stream argument. MLX's own default is the GPU //! when a GPU backend is available (as on Apple Silicon), else the CPU. -use mlxr_sys as sys; +use mlxcore_sys as sys; /// An execution stream bound to a device. pub struct Stream { From 0ef9f71b7fd820dafebc8e785fc95e0d701d292e Mon Sep 17 00:00:00 2001 From: kerthcet Date: Tue, 21 Jul 2026 00:07:39 +0100 Subject: [PATCH 9/9] add assert as precondition check Signed-off-by: kerthcet --- crates/mlxcore/src/array.rs | 19 +++++++++++++++++-- 1 file changed, 17 insertions(+), 2 deletions(-) diff --git a/crates/mlxcore/src/array.rs b/crates/mlxcore/src/array.rs index 998dbf8..ceb39ff 100644 --- a/crates/mlxcore/src/array.rs +++ b/crates/mlxcore/src/array.rs @@ -95,10 +95,18 @@ impl Array { /// The element type `T` selects the accessor at compile time, e.g. /// `a.item::()`. Evaluates the array first. MLX casts the stored dtype /// to `T`. + /// + /// # Panics + /// Panics if the array is not a single-element array (`size() != 1`). pub fn item(&self) -> T { self.eval(); - // SAFETY: the array is evaluated above; `read_item` picks the accessor - // matching `T`. + let size = self.size(); + assert_eq!( + size, 1, + "item() requires a single-element array, but this array has {size} elements" + ); + // SAFETY: `read_item` requires an evaluated, single-element array — both + // ensured above — and picks the accessor matching `T`. unsafe { T::read_item(self.handle) } } @@ -432,6 +440,13 @@ mod tests { assert_eq!(b.item::(), 7); } + #[test] + #[should_panic(expected = "requires a single-element array")] + fn item_on_non_scalar_panics() { + let a = Array::from_slice(&[1.0f32, 2.0, 3.0], &[3]); + let _ = a.item::(); + } + #[test] fn to_vec_is_generic_over_dtype() { let ints = Array::from_slice(&[1i32, 2, 3], &[3]);