# How to Generate Random Numbers Using MLX's Random Module

> Learn to generate random numbers with MLX's random module. Explore PRNG keys, distributions & GPU acceleration for CPU, CUDA, and Metal.

- Repository: [ml-explore/mlx](https://github.com/ml-explore/mlx)
- Tags: how-to-guide
- Published: 2026-06-18

---

**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`](https://github.com/ml-explore/mlx/blob/main/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.

```python
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`](https://github.com/ml-explore/mlx/blob/main/mlx/random.cpp), this supports arbitrary half-precision and full-precision floating-point types.

```python

# 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.

```python

# 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)`.

```python

# 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`](https://github.com/ml-explore/mlx/blob/main/python/tests/test_random.py), supports batches of samples via the `shape` parameter and requires an explicit PRNG key for reproducibility.

```python
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.

```python
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`](https://github.com/ml-explore/mlx/blob/main/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`.

```python

# 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`](https://github.com/ml-explore/mlx/blob/main/python/src/random.cpp) binding layer marshals these arguments to the appropriate backend kernels in [`mlx/random.cpp`](https://github.com/ml-explore/mlx/blob/main/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`](https://github.com/ml-explore/mlx/blob/main/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`](https://github.com/ml-explore/mlx/blob/main/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`](https://github.com/ml-explore/mlx/blob/main/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`](https://github.com/ml-explore/mlx/blob/main/python/tests/test_random.py) and supports generating batches of samples with explicit PRNG keys for reproducibility.