diff --git a/.github/workflows/rust-ci-macos.yaml b/.github/workflows/rust-ci-macos.yaml new file mode 100644 index 0000000..dad36aa --- /dev/null +++ b/.github/workflows/rust-ci-macos.yaml @@ -0,0 +1,37 @@ +name: Rust CI MacOS + +on: + pull_request: + types: + - opened + - synchronize + +env: + CARGO_TERM_COLOR: always + +jobs: + test: + # MLX only builds on Apple Silicon; macos-14+ runners are arm64. + runs-on: macos-15 + steps: + - uses: actions/checkout@v4 + with: + submodules: recursive + + - name: Install Rust toolchain + uses: dtolnay/rust-toolchain@stable + with: + components: rustfmt, clippy + + # Cache the CMake build of mlx-c + MLX, which is the slow part of the build. + - name: Cache cargo & build artifacts + uses: Swatinem/rust-cache@v2 + + - name: Format + run: cargo fmt --all -- --check + + - name: Clippy + run: cargo clippy --all-targets -- -D warnings + + - name: Test + run: make test diff --git a/.gitignore b/.gitignore index e2a094c..e0a2f2b 100644 --- a/.gitignore +++ b/.gitignore @@ -31,3 +31,8 @@ __pycache__ *.pyc .pytest_cache *.tgz + + +# Added by cargo + +/target diff --git a/.gitmodules b/.gitmodules new file mode 100644 index 0000000..b3738cd --- /dev/null +++ b/.gitmodules @@ -0,0 +1,3 @@ +[submodule "third_party/mlx-c"] + path = third_party/mlx-c + url = https://github.com/ml-explore/mlx-c.git diff --git a/Cargo.lock b/Cargo.lock new file mode 100644 index 0000000..9f8661d --- /dev/null +++ b/Cargo.lock @@ -0,0 +1,267 @@ +# This file is automatically @generated by Cargo. +# It is not intended for manual editing. +version = 4 + +[[package]] +name = "aho-corasick" +version = "1.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ddd31a130427c27518df266943a5308ed92d4b226cc639f5a8f1002816174301" +dependencies = [ + "memchr", +] + +[[package]] +name = "bindgen" +version = "0.70.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f49d8fed880d473ea71efb9bf597651e77201bdd4893efe54c9e5d65ae04ce6f" +dependencies = [ + "bitflags", + "cexpr", + "clang-sys", + "itertools", + "log", + "prettyplease", + "proc-macro2", + "quote", + "regex", + "rustc-hash", + "shlex 1.3.0", + "syn", +] + +[[package]] +name = "bitflags" +version = "2.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b588b76d00fde79687d7646a9b5bdf3cc0f655e0bbd080335a95d7e96f3587da" + +[[package]] +name = "cc" +version = "1.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c89588d05638b5b4594a3348a2d6c20277e43a7f5c5202b05cc56888475a47b8" +dependencies = [ + "find-msvc-tools", + "shlex 2.0.1", +] + +[[package]] +name = "cexpr" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6fac387a98bb7c37292057cffc56d62ecb629900026402633ae9160df93a8766" +dependencies = [ + "nom", +] + +[[package]] +name = "cfg-if" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" + +[[package]] +name = "clang-sys" +version = "1.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b023947811758c97c59bf9d1c188fd619ad4718dcaa767947df1cadb14f39f4" +dependencies = [ + "glob", + "libc", + "libloading", +] + +[[package]] +name = "cmake" +version = "0.1.58" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c0f78a02292a74a88ac736019ab962ece0bc380e3f977bf72e376c5d78ff0678" +dependencies = [ + "cc", +] + +[[package]] +name = "either" +version = "1.16.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91622ff5e7162018101f2fea40d6ebf4a78bbe5a49736a2020649edf9693679e" + +[[package]] +name = "find-msvc-tools" +version = "0.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" + +[[package]] +name = "glob" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0cc23270f6e1808e30a928bdc84dea0b9b4136a8bc82338574f23baf47bbd280" + +[[package]] +name = "itertools" +version = "0.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "413ee7dfc52ee1a4949ceeb7dbc8a33f2d6c088194d9f922fb8318faf1f01186" +dependencies = [ + "either", +] + +[[package]] +name = "libc" +version = "0.2.186" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" + +[[package]] +name = "libloading" +version = "0.8.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d7c4b02199fee7c5d21a5ae7d8cfa79a6ef5bb2fc834d6e9058e89c825efdc55" +dependencies = [ + "cfg-if", + "windows-link", +] + +[[package]] +name = "log" +version = "0.4.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ceec5bc11778974d1bcb055b18002eba7f4b3518b6a0081b3af5f21666da9ad" + +[[package]] +name = "memchr" +version = "2.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98" + +[[package]] +name = "minimal-lexical" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a" + +[[package]] +name = "mlx" +version = "0.0.0" +dependencies = [ + "mlx-sys", +] + +[[package]] +name = "mlx-sys" +version = "0.0.0" +dependencies = [ + "bindgen", + "cmake", +] + +[[package]] +name = "nom" +version = "7.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d273983c5a657a70a3e8f2a01329822f3b8c8172b73826411a55751e404a0a4a" +dependencies = [ + "memchr", + "minimal-lexical", +] + +[[package]] +name = "prettyplease" +version = "0.2.37" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "479ca8adacdd7ce8f1fb39ce9ecccbfe93a3f1344b3d0d97f20bc0196208f62b" +dependencies = [ + "proc-macro2", + "syn", +] + +[[package]] +name = "proc-macro2" +version = "1.0.107" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "985e7ec9bb745e6ce6535b544d84d6cd6f7ad8bd711c398938ae983b91a766d9" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "quote" +version = "1.0.47" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1fbf4db142a473a8d80c26bbf18454ed458bf8d26c8219c331daecfdbd079001" +dependencies = [ + "proc-macro2", +] + +[[package]] +name = "regex" +version = "1.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f020237b6c8eed93db2e2cb53c00c60a8e1bc73da7d073199a1180401450218d" +dependencies = [ + "aho-corasick", + "memchr", + "regex-automata", + "regex-syntax", +] + +[[package]] +name = "regex-automata" +version = "0.4.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fcfdb36bda0c880c5931cdc7a2bcdc8ba4556847b9d912bca70bc94708711ad" +dependencies = [ + "aho-corasick", + "memchr", + "regex-syntax", +] + +[[package]] +name = "regex-syntax" +version = "0.8.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4" + +[[package]] +name = "rustc-hash" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08d43f7aa6b08d49f382cde6a7982047c3426db949b1424bc4b7ec9ae12c6ce2" + +[[package]] +name = "shlex" +version = "1.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0fda2ff0d084019ba4d7c6f371c95d8fd75ce3524c3cb8fb653a3023f6323e64" + +[[package]] +name = "shlex" +version = "2.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8fadd59c855ef2080decdef8ff161eb6661b86933c9d82e5ba29dc602a55aba" + +[[package]] +name = "syn" +version = "2.0.119" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "872831b642d1a07999a962a351ed35b955ea2cfc8f3862091e2a240a84f17297" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "unicode-ident" +version = "1.0.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" + +[[package]] +name = "windows-link" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" diff --git a/Cargo.toml b/Cargo.toml new file mode 100644 index 0000000..8333f53 --- /dev/null +++ b/Cargo.toml @@ -0,0 +1,12 @@ +[workspace] +resolver = "2" +members = ["crates/mlx-sys", "crates/mlx"] + +[workspace.package] +edition = "2024" +license = "MIT" +repository = "https://github.com/InftyAI/mlx" +authors = ["kerthcet"] + +[workspace.dependencies] +mlx-sys = { path = "crates/mlx-sys", version = "0.0.0" } diff --git a/Makefile b/Makefile index e69de29..07dbb28 100644 --- a/Makefile +++ b/Makefile @@ -0,0 +1,17 @@ +.PHONY: build test fmt clean + +# Build the workspace (compiles mlx-c + MLX from source on first run). +build: + cargo build + +# Run all tests across the workspace. +test: + cargo test + +# Format all crates. +fmt: + cargo fmt --all + +# Remove build artifacts. +clean: + cargo clean diff --git a/OWNERS b/OWNERS index 4ca4a20..515ec16 100644 --- a/OWNERS +++ b/OWNERS @@ -1,5 +1,5 @@ approvers: - - TBD + - kerthcet reviewers: - - TBD + - kerthcet diff --git a/README.md b/README.md index 4baad06..7bfb93e 100644 --- a/README.md +++ b/README.md @@ -1,3 +1,51 @@ -# template-repo +# MLX -A template repo. +Rust bindings for Apple's [MLX](https://github.com/ml-explore/mlx) framework. + +> **Platform:** Apple Silicon (macOS) only. + +## Architecture + +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 +``` + +MLX itself is C++, which Rust cannot bind to directly. We bind against +[`mlx-c`](https://github.com/ml-explore/mlx-c), Apple's official C API, whose +CMake build pulls in MLX via `FetchContent`. + +## Layout + +- `crates/mlx-sys` — `build.rs` compiles `mlx-c` (+ MLX) with CMake and runs + `bindgen` over `mlx/c/mlx.h`; `src/lib.rs` re-exports the generated bindings. +- `crates/mlx` — safe wrappers (`Array`, ops, …) over `mlx-sys`. +- `third_party/mlx-c` — pinned `mlx-c` submodule (currently `v0.6.0`). + +## Building + +```sh +git submodule update --init --recursive +cargo build +cargo test +``` + +Requirements: a recent Rust toolchain, `cmake`, and Xcode command-line tools. +The first build compiles MLX from source and takes several minutes. + +## Examples + +Runnable examples live in `crates/mlx/examples`: + +```sh +cargo run --example hello +``` + +## Features + +- `metal` (default) — GPU backend via Metal. +- `accelerate` (default) — CPU BLAS via Apple's Accelerate framework. diff --git a/crates/mlx-sys/Cargo.toml b/crates/mlx-sys/Cargo.toml new file mode 100644 index 0000000..40704c8 --- /dev/null +++ b/crates/mlx-sys/Cargo.toml @@ -0,0 +1,23 @@ +[package] +name = "mlx-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. +edition = "2021" +license.workspace = true +repository.workspace = true +authors.workspace = true +description = "Low-level FFI bindings to mlx-c, the C API for Apple's MLX framework." +links = "mlxc" +build = "build.rs" + +[features] +default = ["metal", "accelerate"] +# Build MLX with the Metal (GPU) backend. Apple Silicon only. +metal = [] +# Build MLX against Apple's Accelerate framework for CPU BLAS. +accelerate = [] + +[build-dependencies] +bindgen = "0.70" +cmake = "0.1" diff --git a/crates/mlx-sys/build.rs b/crates/mlx-sys/build.rs new file mode 100644 index 0000000..65d9315 --- /dev/null +++ b/crates/mlx-sys/build.rs @@ -0,0 +1,93 @@ +//! Build script for `mlx-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 +//! `libmlx` come out of a single CMake invocation. +//! 2. Emits the link directives for those static libraries plus the Apple +//! system frameworks MLX depends on. +//! 3. Runs `bindgen` over `mlx/c/mlx.h` (the umbrella header) to produce the +//! raw FFI bindings consumed by `src/lib.rs`. + +use std::env; +use std::path::PathBuf; + +fn main() { + if !cfg!(target_os = "macos") { + panic!("mlx-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 + let mlx_c_dir = manifest_dir + .join("../../third_party/mlx-c") + .canonicalize() + .expect("third_party/mlx-c submodule not found — run `git submodule update --init`"); + + // --- 1. Build mlx-c (+ MLX via FetchContent) with CMake --------------- + let mut cfg = cmake::Config::new(&mlx_c_dir); + cfg.define("BUILD_SHARED_LIBS", "OFF") + .define("MLX_C_BUILD_EXAMPLES", "OFF") + .define("CMAKE_BUILD_TYPE", "Release"); + + if cfg!(feature = "metal") { + cfg.define("MLX_BUILD_METAL", "ON"); + } else { + cfg.define("MLX_BUILD_METAL", "OFF"); + } + if cfg!(feature = "accelerate") { + cfg.define("MLX_BUILD_ACCELERATE", "ON"); + } + + // Build into a fixed dir under `target//` rather than the default + // per-fingerprint `OUT_DIR`. `cargo clippy` fingerprints `mlx-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 + // /. + let out_dir = PathBuf::from(env::var("OUT_DIR").unwrap()); + let profile_dir = out_dir + .ancestors() + .nth(3) + .expect("unexpected OUT_DIR layout"); + cfg.out_dir(profile_dir.join("mlx-c-build")); + + let dst = cfg.build(); + + // --- 2. Link directives ---------------------------------------------- + // CMake installs archives under /lib. + println!("cargo:rustc-link-search=native={}/lib", dst.display()); + println!("cargo:rustc-link-lib=static=mlxc"); + println!("cargo:rustc-link-lib=static=mlx"); + + // MLX is C++; pull in the C++ standard library. + println!("cargo:rustc-link-lib=dylib=c++"); + + // Apple system frameworks MLX links against. + for framework in ["Foundation", "Metal", "QuartzCore", "Accelerate"] { + println!("cargo:rustc-link-lib=framework={framework}"); + } + + // --- 3. Generate bindings -------------------------------------------- + let header = mlx_c_dir.join("mlx/c/mlx.h"); + println!("cargo:rerun-if-changed={}", header.display()); + + let bindings = bindgen::Builder::default() + .header(header.to_string_lossy()) + // mlx-c headers `#include "mlx/c/..."` relative to the repo root. + .clang_arg(format!("-I{}", mlx_c_dir.display())) + .allowlist_function("mlx_.*") + .allowlist_type("mlx_.*") + .allowlist_var("MLX_.*") + .prepend_enum_name(false) + .parse_callbacks(Box::new(bindgen::CargoCallbacks::new())) + .generate() + .expect("failed to generate mlx-c bindings"); + + // bindings.rs stays in the real OUT_DIR — it's cheap to regenerate and + // `src/lib.rs` includes it from there. + bindings + .write_to_file(out_dir.join("bindings.rs")) + .expect("failed to write bindings.rs"); +} diff --git a/crates/mlx-sys/src/lib.rs b/crates/mlx-sys/src/lib.rs new file mode 100644 index 0000000..5269012 --- /dev/null +++ b/crates/mlx-sys/src/lib.rs @@ -0,0 +1,12 @@ +//! Raw, unsafe FFI bindings to [`mlx-c`](https://github.com/ml-explore/mlx-c), +//! the C API for Apple's MLX framework. +//! +//! 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. +#![allow(non_upper_case_globals)] +#![allow(non_camel_case_types)] +#![allow(non_snake_case)] +#![allow(dead_code)] + +include!(concat!(env!("OUT_DIR"), "/bindings.rs")); diff --git a/crates/mlx/Cargo.toml b/crates/mlx/Cargo.toml new file mode 100644 index 0000000..bebe917 --- /dev/null +++ b/crates/mlx/Cargo.toml @@ -0,0 +1,16 @@ +[package] +name = "mlx" +version = "0.0.0" +edition.workspace = true +license.workspace = true +repository.workspace = true +authors.workspace = true +description = "Safe, idiomatic Rust bindings for Apple's MLX array framework." + +[features] +default = ["metal", "accelerate"] +metal = ["mlx-sys/metal"] +accelerate = ["mlx-sys/accelerate"] + +[dependencies] +mlx-sys.workspace = true diff --git a/crates/mlx/examples/hello.rs b/crates/mlx/examples/hello.rs new file mode 100644 index 0000000..a74a468 --- /dev/null +++ b/crates/mlx/examples/hello.rs @@ -0,0 +1,18 @@ +//! A minimal MLX example. +//! +//! Run with: +//! ```sh +//! cargo run --example hello +//! ``` + +use mlx::Array; + +fn main() { + println!("MLX version: {}", mlx::version()); + + let a = Array::from_slice(&[1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]); + println!("shape: {:?}", a.shape()); + println!("size: {}", a.size()); + println!("ndim: {}", a.ndim()); + println!("array:\n{a:?}"); +} diff --git a/crates/mlx/src/array.rs b/crates/mlx/src/array.rs new file mode 100644 index 0000000..d09b916 --- /dev/null +++ b/crates/mlx/src/array.rs @@ -0,0 +1,144 @@ +//! A safe wrapper around `mlx_array`. + +use std::ffi::CStr; +use std::fmt; + +use mlx_sys as sys; + +use crate::dtype::ArrayElement; + +/// An N-dimensional MLX array. +/// +/// Owns the underlying `mlx_array` handle and frees it on drop. +pub struct Array { + handle: sys::mlx_array, +} + +impl Array { + /// Creates an `Array` from a raw handle, taking ownership of it. + /// + /// # Safety + /// `handle` must be a valid `mlx_array` that is not freed elsewhere. + pub(crate) unsafe fn from_raw(handle: sys::mlx_array) -> Self { + Self { handle } + } + + /// 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 + } + + /// Builds an array from a slice of values with the given shape. + /// + /// The MLX dtype is chosen from the element type `T` at compile time (e.g. + /// `&[f32]` produces a `float32` array, `&[i32]` an `int32` array). + /// + /// # Panics + /// Panics if `data.len()` does not equal the product of `shape`. + pub fn from_slice(data: &[T], shape: &[i32]) -> Self { + let expected: i64 = shape.iter().map(|&d| d as i64).product(); + assert_eq!( + data.len() as i64, + expected, + "data length {} does not match shape product {expected}", + data.len() + ); + // SAFETY: pointers/len are valid for the duration of the call; mlx + // copies the data into its own buffer. + let handle = unsafe { + sys::mlx_array_new_data( + data.as_ptr() as *const _, + shape.as_ptr(), + shape.len() as i32, + T::DTYPE, + ) + }; + unsafe { Self::from_raw(handle) } + } + + /// Total number of elements. + pub fn size(&self) -> usize { + // SAFETY: handle is valid for the lifetime of `self`. + unsafe { sys::mlx_array_size(self.handle) } + } + + /// Number of dimensions. + pub fn ndim(&self) -> usize { + unsafe { sys::mlx_array_ndim(self.handle) } + } + + /// Shape of the array. + pub fn shape(&self) -> Vec { + let ndim = self.ndim(); + // SAFETY: mlx guarantees the returned pointer is valid for `ndim` ints. + let ptr = unsafe { sys::mlx_array_shape(self.handle) }; + (0..ndim).map(|i| unsafe { *ptr.add(i) }).collect() + } +} + +impl fmt::Debug for Array { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + // 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) }; + let cstr = unsafe { CStr::from_ptr(sys::mlx_string_data(s)) }; + let out = write!(f, "{}", cstr.to_string_lossy()); + unsafe { sys::mlx_string_free(s) }; + out + } +} + +impl Drop for Array { + fn drop(&mut self) { + // SAFETY: `handle` was created by mlx and is owned solely by `self`. + unsafe { + sys::mlx_array_free(self.handle); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn from_slice_reports_shape_size_ndim() { + let a = Array::from_slice(&[1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]); + assert_eq!(a.size(), 6); + assert_eq!(a.ndim(), 2); + assert_eq!(a.shape(), vec![2, 3]); + } + + #[test] + fn element_type_selects_dtype() { + // Both element types build valid arrays; the dtype is carried by `T`. + let floats = Array::from_slice(&[1.0f32, 2.0], &[2]); + assert_eq!(floats.shape(), vec![2]); + let ints = Array::from_slice(&[1i32, 2, 3], &[3]); + assert_eq!(ints.shape(), vec![3]); + } + + #[test] + fn scalar_array_is_zero_dim() { + let a = Array::from_slice(&[42.0f32], &[]); + assert_eq!(a.ndim(), 0); + assert_eq!(a.size(), 1); + assert!(a.shape().is_empty()); + } + + #[test] + #[should_panic(expected = "does not match shape product")] + fn mismatched_len_and_shape_panics() { + // 5 elements cannot fill a 2x3 (=6) array. + let _ = Array::from_slice(&[1.0f32, 2.0, 3.0, 4.0, 5.0], &[2, 3]); + } + + #[test] + fn debug_renders_array_contents() { + let a = Array::from_slice(&[1.0f32, 2.0], &[2]); + let s = format!("{a:?}"); + assert!(s.contains("array"), "unexpected debug output: {s}"); + } +} diff --git a/crates/mlx/src/dtype.rs b/crates/mlx/src/dtype.rs new file mode 100644 index 0000000..c359c4f --- /dev/null +++ b/crates/mlx/src/dtype.rs @@ -0,0 +1,58 @@ +//! Mapping between Rust primitive types and MLX data types. + +use mlx_sys as sys; + +mod sealed { + pub trait 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 { + /// The MLX dtype corresponding to this Rust type. + const DTYPE: sys::mlx_dtype; +} + +macro_rules! impl_array_element { + ($($rust:ty => $dtype:expr),* $(,)?) => { + $( + impl sealed::Sealed for $rust {} + impl ArrayElement for $rust { + const DTYPE: sys::mlx_dtype = $dtype; + } + )* + }; +} + +// Only the MLX dtypes with a native Rust primitive are mapped here. Types +// 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, +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn maps_rust_types_to_expected_dtypes() { + assert_eq!(::DTYPE, sys::MLX_FLOAT32); + assert_eq!(::DTYPE, sys::MLX_FLOAT64); + assert_eq!(::DTYPE, sys::MLX_INT32); + assert_eq!(::DTYPE, sys::MLX_UINT8); + assert_eq!(::DTYPE, sys::MLX_BOOL); + } +} diff --git a/crates/mlx/src/lib.rs b/crates/mlx/src/lib.rs new file mode 100644 index 0000000..1b22be7 --- /dev/null +++ b/crates/mlx/src/lib.rs @@ -0,0 +1,35 @@ +//! 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. +//! +//! This crate is Apple Silicon (macOS) only. + +mod array; +mod dtype; + +pub use array::Array; +pub use dtype::ArrayElement; + +/// Returns the version string of the underlying MLX library. +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)) + .to_string_lossy() + .into_owned(); + mlx_sys::mlx_string_free(s); + v + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn reports_version() { + assert!(!version().is_empty()); + } +} diff --git a/third_party/mlx-c b/third_party/mlx-c new file mode 160000 index 0000000..0726ca9 --- /dev/null +++ b/third_party/mlx-c @@ -0,0 +1 @@ +Subproject commit 0726ca922fc902c4c61ef9c27d94132be418e945