Skip to content

CAX: Cellular Automata Accelerated in JAX

logo
PyPI - Python Version PyPI - Version Paper X URL

CAX is a high-performance and flexible open-source library designed to accelerate artificial life research — cellular automata, particle systems, and other self-organizing complex systems, all in JAX. 🧬

Overview 🔎

Are you interested in emergence, self-organization, or open-endedness? Whether you're a researcher or just curious about the fascinating world of artificial life, CAX is your digital lab! 🔬

Designed for speed and flexibility, CAX allows you to easily experiment with self-organizing behaviors and emergent phenomena. 🧑‍🔬

Get started here Colab

Why CAX? 💡

CAX supports discrete and continuous systems, including neural cellular automata, across any number of dimensions. Beyond traditional cellular automata, it also handles particle systems and more, all unified under a single, intuitive API.

Rich 🎨

CAX provides a comprehensive collection of 25+ ready-to-use systems. From simulating one-dimensional elementary cellular automata to training three-dimensional self-autoencoding neural cellular automata, or even creating beautiful Lenia simulations, CAX provides a versatile platform for exploring the rich world of self-organizing systems.

Flexible 🧩

CAX makes it easy to extend existing systems or build custom ones from scratch for endless experimentation and discovery. Design your own experiments to probe the boundaries of artificial open-ended evolution and emergent complexity.

Fast 🚀

CAX is built on top of the JAX/Flax ecosystem for speed and scalability. The library benefits from vectorization and parallelization on various hardware accelerators such as CPU, GPU, and TPU. This allows you to scale your experiments from small prototypes to massive simulations with minimal code changes.

Tested & Documented 📚

The library is thoroughly tested and documented with numerous examples to get you started! Our comprehensive guides walk you through everything from basic cellular automata to advanced neural implementations.

Examples 📓

# Example Reference Colab
1 Elementary Cellular Automata Wolfram (2002) Colab
2 Conway's Game of Life Gardner (1970) Colab
3 Langton's Ant Langton (1986) Colab
4 Abelian Sandpile Bak et al. (1987) Colab
5 Lenia Chan (2020) Colab
6 Flow Lenia Plantec et al. (2022) Colab
7 Particle Lenia Mordvintsev et al. (2022) Colab
8 Reaction-Diffusion Gray & Scott (1984) Colab
9 Particle Life Mohr (2018) Colab
10 Boids Reynolds (1987) Colab
11 Neural Particle Automata Kim et al. (2026) Colab
12 Growing Neural Cellular Automata Mordvintsev et al. (2020) Colab
13 Growing Conditional NCA Sudhakaran et al. (2022) Colab
14 Growing Unsupervised NCA Palm et al. (2021) Colab
15 Diffusing Neural Cellular Automata Faldor et al. (2024) Colab
16 Self-classifying MNIST Digits Randazzo et al. (2020) Colab
17 Self-autoencoding MNIST Digits Faldor et al. (2024) Colab
18 Texture Neural Cellular Automata Niklasson et al. (2021) Colab
19 1D-ARC Neural Cellular Automata Faldor et al. (2024) Colab
20 Attention-based Neural Cellular Automata Tesfaldet et al. (2022) Colab
21 Isotropic Neural Cellular Automata Mordvintsev et al. (2022) Colab
22 Differentiable Logic Cellular Automata Miotti et al. (2025) Colab
23 Variational Autoencoder Kingma & Welling (2013) Colab
24 Recurrent Residual CNN Faldor et al. (2024) Colab
25 Growing NCA with Evolution Strategies Faldor et al. (2024) Colab
26 Growing NCA with Reinforcement Learning Faldor et al. (2024) Colab
27 Leniabreeder Faldor & Cully (2024) Colab
28 Gradient Descent in Lenia Hamon et al. (2024) Colab
29 Lenia Gradients in Depth Faldor et al. (2024) Colab

Getting Started 🚦

Here, you can see the basic CAX API usage with Conway's Game of Life:

import jax
import jax.numpy as jnp

from cax.cs.life import Life

seed = 0

num_steps = 128
spatial_dims = (32, 32)
channel_size = 1
rule_golly = "B3/S23"  # Conway's Game of Life

key = jax.random.key(seed)

birth, survival = Life.birth_survival_from_string(rule_golly)
cs = Life(birth=birth, survival=survival)

state_init = jax.random.bernoulli(
    key, p=0.5, shape=(*spatial_dims, channel_size)
).astype(jnp.float32)
state_final, states = cs(state_init, num_steps=num_steps, return_states=True)

For a more detailed overview, get started with this notebook Colab

Installation ⚙️

You will need Python 3.12 or later, and a working JAX installation installed in a virtual environment.

Then, install CAX from PyPi with uv:

uv pip install cax

or with pip:

pip install cax

Citing CAX 📝

If you use CAX in your research, please cite the following paper:

@inproceedings{cax,
    title = {{CAX}: {Cellular} {Automata} {Accelerated} in {JAX}},
    volume = {2025},
    url = {https://proceedings.iclr.cc/paper_files/paper/2025/file/19206a6ed5ed0aaeed440448dfc5cf7e-Paper-Conference.pdf},
    booktitle = {International {Conference} on {Representation} {Learning}},
    author = {Faldor, Maxence and Cully, Antoine},
    editor = {Yue, Y. and Garg, A. and Peng, N. and Sha, F. and Yu, R.},
    year = {2025},
    pages = {8947--8960},
    keywords = {artificial life, emergence, self-organization, open-endedness, cellular automata, neural cellular automata},
}

Contributing 👷

Contributions are welcome! If you find a bug or are missing your favorite self-organizing system, please open an issue or submit a pull request following our contribution guidelines 🤗.