Notes on MLX
Collection of my notes from reading the MLX docs.
Memory Bandwidth
- Apple Silicon unified memory
- M1: 68 GB/s
- M2/M3: 102 GB/s
- M4: 120 GB/s
- M4 Pro: 273 GB/s
- M4 Max: 410 - 546 GB/s
- M5: 153 GB/s
- M5 Pro: 307 GB/s
- M5 Max: 460 - 614 GB/s
- M5 Ultra: 1228.8 GB/s
- M6: 153 - 170 GB/s
- System memory
- DDR5: 32 - 70 GB/s
- DDR4: 12.8 - 25.6 GB/s
- DDR3: 6.4 - 17 GB/s
- Nvidia GPU VRAM
- GDDR7 on 5090: 1.79 TB/s
- GDDR6X on 4090: 1008 GB/s
Roughly then, a M5 Ultra chip has the slightly more memory bandwidth than a RTX 4090 but less than a RTX 5090. It is however vastly greater in capacity being 96/256 GB compared to either 24 or 32 GB for the GPUs.
MLX uses Lazy Evaluation
Operations create a compute graph that is evaluated when values are accessed for the first time.
def fun(x):
a = fun1(x)
b = expensive_fun(a)
return a, b
y, _ = fun(x)
Here, expensive_fun is never computed.
Unified Memory
Both the CPU and the GPU share the same unified memory, thus data does not need to be moved across devices.
Instead the stream parameter is used to control the device that is used to execute the operation.
import mlx.core as mx
a = mx.random.normal((100,))
b = mx.random.normal((100,))
mx.add(a, b, stream=mx.cpu)
mx.add(a, b, stream=mx.gpu)
Dependency is automatically resolved by the scheduler based on the compute graph.
This can be exploited for scheduling operations with higher arithmetic intensity on the GPU (eg. matmuls!) while lower ones can be scheduled on the CPU (eg. pointwise operations).
Compilation
Compiling computation graphs with mx.compile produces smaller graphs. This is achieved by merging common operations and fusing.
Some gotchas with compilation:
- Compiled functions are traced with placeholder inputs. Eval on arrays inside functions would thus not work.
- Compiled functions should be side effect free.
Numerical Precision
MLX can automatically run operations at reduced precision on supported hardware.
Inputs and outputs to these operations however remain at the original intended precision.
This can be disabled by setting the environment variable MLX_ENABLE_TF32=0.
Distributed Communication
Supported backeds are: MPI, RING, JACCL (needs fully connected topologies), NCCL. mlx.launch is used to launch distributed jobs across a fleet.