MLX: An array framework for Apple Silicon
github.com
github.com
Flax+Jax+OpenXLA seems to finally be building some momentum so when a big player launches yet another competitor, the justification for it would be a good thing to see.
What was “not good enough” with Jax? Why did it make sense to put this human time and energy there instead of doubling down on Flax/Jax/OpenXLA? How will this move the needle at all against Nvidia?
I guess the fact that Google is pushing OpenXLA is making the other giants not want to truly lean in?
I don’t want 10 competing “choices”. I want one clear, open, competitor to Cuda that works on all the competing hardware.
x[0] = 10
And instead I have to do: y = x.at[0].set(10)
Of course this has advantages, and I know it sounds lame, but as someone whose brain works in numpy, this is really offputting. arr[1000000] = 1
has to clone the entire array if you want it to be pure, leading to very unpredictable performance. There are also some algorithms that are straight up impossible to (efficiently) implement without mutability. Often, it's exactly those algorithms that are hard to optimize for optimizers like JAX.Specifically in JAX, code that is slow due to copying will often be optimized into mutable code before running for performance reasons. But because JAX still has the gurantees of no mutability, it can do many optimizations such as caching or dead-code elimination.
You are correct about in-depth mutations and resulting complications, but that only strengthens my assertion: immutable should be default, but not compulsory (because sometimes you absolutely need them). And mutability doesn't preclude caching or dead-code elimination; you just have to be more careful. Often it's the case that you can convert a mutable code into an immutable form only for the purpose of analysis, which is definitely harder than an immutable code in the first place but not impossible. Scalar compilers have used SSA---an immutable description for mutable programs---for a long time after all.
for ex, a simple loop in JAX:
def solve(i, v): return i+v
x = jax.lax.fori_loop(0, 5, solve, 10)Many keep forgeting that CUDA for years is a polyglot platform, C, C++, Fortran, plus anything PTX, some of which also target OpenCL, meaning Haskell, Java, C#, Julia, Futhark, or Python bindings.
Then there are the libraries, and GPGPU graphical debugging tools.
By the way, Modular just announced partnerships with AWS and NVidia for Mojo and related tooling.
Running non-trivial ML workloads on the edge has been on my wishlist for years and it sounds like Apple has just the thing.
Other frameworks have (non-python) solutions for mobile in place, e.g. tflite, libtorch, and onnxruntime.
The quick start guide has an overview.
> The Python API closely follows NumPy with a few exceptions. MLX also has a fully featured C++ API which closely follows the Python API.
The main differences between MLX and NumPy are:
Composable function transformations: MLX has composable function transformations for automatic differentiation, automatic vectorization, and computation graph optimization.
Lazy computation: Computations in MLX are lazy. Arrays are only materialized when needed.
Multi-device: Operations can run on any of the supported devices (CPU, GPU, …)
The design of MLX is inspired by frameworks like PyTorch, Jax, and ArrayFire. A noteable difference from these frameworks and MLX is the unified memory model. Arrays in MLX live in shared memory. Operations on MLX arrays can be performed on any of the supported device types without performing data copies. Currently supported device types are the CPU and GPU.AFAIK the only way to leverage Apple Neural Engine (and get the best performance) is to use CoreML. The only documented way to use CoreML is via coremltools, which takes a trace of a PyTorch model and attempts to translate it into a protobuf graph understood by CoreML.
This process often fails and requires model changes, or worse "succeeds" but gives you the wrong output when you run the model. Additionally, you have to play detective to figure out why some operations run on the CPU, or GPU instead of ANE.
It's exciting to see more tools like this for working with tensor-like objects, but I really wish Apple would make porting custom models in a high performance manner easier.
That said I've had good success with onnxruntime recently [0].
[0] https://onnxruntime.ai/docs/execution-providers/CoreML-Execu...
Thus other platforms can simply take this backend-code and integrate it. (Pytorch basically did that already with Apple's help).
The process becomes particularly frustrating when the model appears to convert successfully, but then fails to produce any output or loses layers entirely. Additionally, the debug information provided by the conversion tool isn't very helpful, adding to the challenge.
As an iOS developer with no prior experience in Python, I found myself in a unique position. I needed to build a custom model for one of my keyboard apps to handle tasks like spellchecking, grammar correction, next-word prediction, and autocompletion. This necessity pushed me to learn Python and PyTorch. After mastering these, I then had to convert my knowledge back to Swift and CoreML.
Ideally, I would have preferred to build my model directly in Swift & CoreML, but the current tools and resources for this approach are limited. This limitation is particularly evident in terms of the ease of use and flexibility that Python and PyTorch offer.
That being said, going to spend some time getting familiar with it by reimplementing Llama-2 and trying to make it fast: https://github.com/jbarrow/mlxllama
Otherwise I wonder if this really finds too much adoption. You don't want to lockin yourself when you have all those other choices. (Ok, to be fair, as it is discussed here, PyTorch etc might not work optimal yet on Apple Silicon, but I guess this is just a matter of time.)
I already have this nice, powerful, and expensive hardware sitting on my desk. Why not make the most optimal use of it? I can worry about lock-in once it becomes inadequate for the task.
> PyTorch etc might not work optimal yet on Apple Silicon
So now we can take MLX apart and see how we can use it to improve PyTorch.
The README mentions unified memory, but what stops other frameworks from modeling copies as no-ops? I wonder if MLX makes larger architectural decisions based on GPU CPU communication being cheap.
[0] https://x.com/awnihannun/status/1732184443451019431?s=46&t=O...
Literally “thumbs up” to “does it use GPU?”? Uh…
The links on that thread are the same links from the top of the GitHub repo.
I mean, here’s a nitter link https://nitter.net/awnihannun/status/1732184443451019431#m for anyone else who’s interested, but the info on x seems to be a nothing that isn’t already on the GitHub page.
The GitHub repo seems to be more useful and informative.
That is how the kidz do deep dives these days.
There are few limitations left when compared with other backends. Instead of using 'cuda' device, one simply uses 'MPS' as device.
What remains is: the optimizations Pytorch provides (especially compile() with 2.1) focus on cuda and it's historic restrictions that result from CUDA being _not_ unified memory, and lots of energy goes into developing architectural work-arounds in order to limit the copying between graphics HW and CPU memory, resulting in proprietary compilers (like triton) that move parts of the python code into proprietary hardware.
Apple's unified memory would make all of those super complicated architectural workarounds mostly unnecessary (which they demonstrate with their project).
Getting current D/L platforms to support both paradigms (unified/non unified) will be a lot of work. One possible avenue is the MLIR project currently leveraged by Mojo.
Not sure what you're referring to, the link I provided shows how to use the "mps" backend / device from the official PyTorch release.
> lots of energy goes into developing architectural work-arounds in order to limit the copying between graphics HW and CPU memory
Does this remark apply to PyTorch running on NVidia's platforms with unified memory like the Jetsons?
einsum seems like a reasonable thing to request, but it's hard to be performant across the entire surface exposed by the operation.
https://github.com/pytorch/pytorch/issues/77764
NVIDIA's moat is not just in providing BLAS++ operations, but extending this to a wider range of cuSPARSE, cuSOLVE, cuTENSOR, etc. Without these, it feels like Apple is just trying to play catch up with whatever is popular and unsupported...
Probably reading into this too much, but is this hinting at future Neural Engine support?
It’d be nice to access that without CoreML.
Does this mean its from Apple directly (as part of the machine-learning engineering team)?