How to Generate Random Numbers Using MLX's Random Module
MLX provides a deterministic, hardware-accelerated random number generation API under mlx.random that supports explicit PRNG keys, multiple statistical distributions, and seamless execution across CPU, CUDA, and Metal backends.
The ml-explore/mlx repository implements a NumPy-compatible random module that separates state management from tensor generation using pure functional PRNG keys. By leveraging explicit keys and device-agnostic sampling functions, you can generate reproducible random numbers on Apple Silicon, NVIDIA GPUs, and standard CPUs without modifying source code.
PRNG Keys and Global Seeding in MLX
MLX adopts a functional approach to randomness. Instead of modifying hidden global state, you explicitly create PRNG keys using mx.random.key(seed) or set a global default with mx.random.seed(seed). As shown in python/tests/test_random.py, the seed function initializes the default key used when no explicit key is provided, while key creates specific keys for fine-grained control.
import mlx as mx
# Set global seed for implicit key usage
mx.random.seed(123)
# Create explicit key for reproducible streams
key = mx.random.key(42)
All subsequent random operations can accept an optional key parameter to override the default, ensuring deterministic behavior across runs.
Generating Uniform and Normal Random Numbers
The mlx.random module provides vectorized samplers for common distributions. These functions accept a shape tuple and optional dtype and stream parameters to control output tensor properties.
Uniform distribution samples are generated via mx.random.uniform(), which draws values from a uniform range. According to the implementation in mlx/random.cpp, this supports arbitrary half-precision and full-precision floating-point types.
# Uniform floats in [0, 1)
x = mx.random.uniform(shape=(5, 3))
# Bounded uniform with bfloat16 precision
x = mx.random.uniform(shape=(1000,), low=-1, high=5, dtype=mx.bfloat16)
Normal and Laplace distributions are similarly available via mx.random.normal() and mx.random.laplace(). The Metal and CUDA kernels in mlx/backend/metal/kernels/random.metal and mlx/backend/cuda/random.cu accelerate these operations on Apple Silicon and NVIDIA hardware.
# Standard normal distribution
noise = mx.random.normal(shape=(2, 4))
# Laplace with location=0, scale=1
samples = mx.random.laplace(shape=(3, 3))
Sampling Integers and Multivariate Distributions
For discrete sampling, mx.random.randint() generates random integers within a specified half-open interval [low, high).
# Random integers in [0, 10)
ints = mx.random.randint(low=0, high=10, shape=(4, 4))
Multivariate normal sampling requires explicit mean and covariance arrays. The mx.random.multivariate_normal() function, tested in python/tests/test_random.py, supports batches of samples via the shape parameter and requires an explicit PRNG key for reproducibility.
mean = mx.array([0.0, 0.0])
cov = mx.array([[1.0, 0.5],
[0.5, 2.0]])
# Generate 1000 samples
samples = mx.random.multivariate_normal(mean, cov, shape=(1000,), key=mx.random.key(0))
Reproducible Parallelism with Key Splitting
To create independent random streams for parallel workloads, use mx.random.split(). This function derives multiple statistically independent keys from a single parent key, enabling reproducible parallelism without sequence overlaps.
master_key = mx.random.key(0)
# Split into 4 independent keys
sub_keys = mx.random.split(master_key, 4)
# Use in parallel streams
part0 = mx.random.uniform(shape=(10,), key=sub_keys[0])
part1 = mx.random.normal(shape=(10,), key=sub_keys[1])
As implemented in the Python bindings in python/src/random.cpp, splitting preserves the deterministic properties of the underlying PRNG algorithm used across all backends.
Device and Data Type Control
Every random sampler accepts dtype and stream arguments to control the output tensor's data type and physical device placement. The stream parameter accepts device objects like mx.cpu, mx.gpu, or mx.cuda, while dtype supports mx.float32, mx.float16, and mx.bfloat16.
# Generate directly on GPU
gpu_tensor = mx.random.uniform(shape=(256, 256), stream=mx.gpu)
# Use half-precision floats
small_float = mx.random.normal(shape=(100,), dtype=mx.float16)
The python/src/random.cpp binding layer marshals these arguments to the appropriate backend kernels in mlx/random.cpp or the GPU-specific implementations.
Backend Implementation Architecture
The random module's performance stems from hardware-specific kernels. The core logic resides in mlx/random.cpp for CPU execution, with GPU acceleration provided by mlx/backend/cuda/random.cu for NVIDIA devices and mlx/backend/metal/kernels/random.metal for Apple Silicon. The Python API exposed in python/src/random.cpp provides a unified interface that dispatches to these backends based on the requested stream, ensuring consistent behavior across platforms.
Summary
- PRNG Keys: Use
mx.random.key()for explicit state ormx.random.seed()for global defaults to ensure reproducibility. - Distribution Samplers: Generate uniform, normal, Laplace, and integer random numbers via
uniform(),normal(),laplace(), andrandint(). - Multivariate Support: Sample from multivariate normal distributions using
multivariate_normal()with explicit mean and covariance matrices. - Parallel Streams: Split keys with
mx.random.split()to create independent random sequences for parallel computation. - Hardware Control: Target specific devices (CPU, CUDA, Metal) and data types using the
streamanddtypeparameters. - Implementation: Native performance is delivered through optimized C++, CUDA, and Metal kernels.
Frequently Asked Questions
How do I set a global random seed in MLX?
Call mx.random.seed(seed_value) to initialize the global default PRNG key. Subsequent calls to random functions that do not specify a key argument will use this seeded state, as demonstrated in python/tests/test_random.py.
What is the difference between mx.random.key and mx.random.seed?
mx.random.key(seed) creates a new PRNG key object that you pass explicitly to random functions, enabling fine-grained control. mx.random.seed(seed) sets a global default key that is used implicitly when no key is provided, similar to traditional global random state management.
How can I generate random numbers on a specific device in MLX?
Pass the stream argument to any random function, such as mx.random.uniform(shape=(10,), stream=mx.gpu) for GPU generation. The underlying kernels in mlx/backend/cuda/random.cu or mlx/backend/metal/kernels/random.metal handle the device-specific execution.
Does MLX support multivariate normal distributions?
Yes, use mx.random.multivariate_normal(mean, cov, shape=..., key=...) where mean is a 1-D array and cov is a 2-D square matrix. This function is tested in python/tests/test_random.py and supports generating batches of samples with explicit PRNG keys for reproducibility.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →