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 or mx.random.seed() for global defaults to ensure reproducibility.
  • Distribution Samplers: Generate uniform, normal, Laplace, and integer random numbers via uniform(), normal(), laplace(), and randint().
  • 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 stream and dtype parameters.
  • 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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →