Library / SDK
google-deepmind/mctx avatar
google-deepmind/mctx

Mctx: JAX-native Monte Carlo tree search for MuZero-style agents

Monte Carlo tree search in JAX

2,667 stars217 forksPythonApache-2.0

At a glance

What is it?
Mctx is a DeepMind library that implements MCTS inside JAX, so search runs on the same accelerators and under the same jit as the rest of a Python RL stack. It is for researchers with a learned model to plug in, not for game players who want a ready-made bot.
Who is it for?
Adopt Mctx if you already train in JAX and can express your environment as a recurrent_fn plus a RootFnOutput, and if you want search to run batched and jitted rather than in a Python loop. Do not adopt it if you need a turnkey opponent for a board game or you cannot supply learned priors and values, since Mctx has no environment of its own.
Can I use it commercially?
Yes. Apache-2.0 is a permissive licence: you can use, modify and sell software built on it, as long as you keep its copyright and licence notices.
Is it still maintained?
Yes. The repository last received commits 20 days ago.
What is it written in?
Mainly Python, according to GitHub's language statistics.

Answers come from the project's GitHub data, last synced on September 24, 2026, and from our analysis. They are not legal advice.

Editorial analysis

What Mctx solves, and who it is actually for

Mcts is easy to write and hard to write fast. A conventional implementation walks a tree node by node in Python, which means the search speed is set by the interpreter rather than by the hardware. The fast versions are typically written in C++, and the README is explicit that this "can come at the expense of usability and hackability, especially for researchers that are not familiar with C++."

Mctx takes the other route. It implements MCTS natively in JAX, so the tree search is a JAX computation that can be JIT-compiled and dispatched to an accelerator. The README states that search algorithms "are defined for and operate on batches of inputs, in parallel," which is what lets the library work with large learned environment models parameterized by deep neural networks.

The audience follows from that. This is a library for people doing search-based reinforcement learning research in Python, who have a learned model and want to try variations on the search itself. It is not a game engine, it does not ship an environment, and it will not play anything for you until you supply the model. If you want a chess or Go opponent out of the box, you are looking at the wrong project.

The recurrent_fn contract that defines the whole architecture

Mctx does not know what your environment is. Everything it needs arrives through two objects described in the README.

The root state is described by a RootFnOutput, which carries prior_logits from a policy network, the estimated value of the root state, and an embedding that represents that state for the environment model. The dynamics are described by a recurrent_fn, called as recurrent_fn(params, rng_key, action, embedding), returning a tuple of a RecurrentFnOutput and the next embedding. The RecurrentFnOutput holds the reward and discount for the transition, plus prior_logits and value for the new state.

That is the entire interface. Search quality is therefore bounded by the model you hand in: if the priors are poor, the tree expands in the wrong direction, and if the value estimates are wrong, the backup propagates the error. The README is candid about the second case in the context of Gumbel MuZero, which "guarantees a policy improvement if the action values are correctly evaluated." The conditional is doing real work there.

Above that contract sit three levels of API: a low-level generic search function, and two concrete policies, muzero_policy and gumbel_muzero_policy.

Installing Mctx and running a first batched search

The released version comes from PyPI. The README gives this command:

bash
pip install mctx

There is also a development install straight from the repository, which the README documents as:

bash
pip install git+https://github.com/google-deepmind/mctx.git

Dependencies are declared in pyproject.toml: chex>=0.1.91, jax>=0.7.0 and jaxlib>=0.7.0. Python 3.11 or newer is required. If your existing JAX pin is older than 0.7.0, pip will try to move it, and that can disturb the rest of your stack. Check the pin before installing rather than after.

The optional extras matter for the examples. The test extra adds absl-py and numpy; the examples extra adds absl-py and pygraphviz. The visualization demo draws trees, so it needs pygraphviz, which in turn needs Graphviz installed on the system.

Once installed, the call pattern from the README's quickstart is a single function:

python
policy_output = mctx.gumbel_muzero_policy(params, rng_key, root, recurrent_fn,
                                          num_simulations=32)

policy_output.action is the action the search proposes, which you pass to your environment. policy_output.action_weights are targets usable to train the policy probabilities. The README recommends gumbel_muzero_policy over the plain MuZero policy, and points at examples/policy_improvement_demo.py as the demonstration of the improvement guarantee.

Where Mctx is the wrong tool

The first limitation is structural: Mctx has no environment. There is no board, no state transition, no reward function, and no terminal condition in the library. All of it lives in your recurrent_fn. If you cannot express your environment as a pure function of (params, rng_key, action, embedding), Mctx cannot search it.

The second is the value-estimation dependency. The policy improvement result the README cites is conditional on action values being correctly evaluated. In early training they are not, and nothing in the library checks that for you. A search over bad values produces a confident bad policy, and the failure looks like a search problem when it is a learning problem.

The third is scale. Because the implementation is batched and jitted, the natural unit of work is a batch of searches, not one. A single small search, run once, gains little from JAX and pays the compilation cost. The library's own framing is about making the most of accelerators with large learned models; if your model is a lookup table, a plain Python MCTS will be simpler to debug and probably no slower.

Finally, the project describes itself as Beta in its own classifiers. Treat the API as something that can still shift between releases.

How it compares with a hand-written Python MCTS

The obvious alternative is to write MCTS yourself in Python. The difference is not features, it is where the loop lives. A hand-written tree search is a Python while-loop over node selections and expansions, and each node visit costs interpreter time. That is fine for small trees and for teaching, and it is far easier to step through in a debugger.

Mctx inverts the trade. The search becomes a JAX program: batched, jittable, and able to run on the same device as your network. What you give up is the ability to inspect a mutable tree object mid-search. Debugging moves from breakpoints to shape errors and traced values.

There is also a family of compiled C++ MCTS implementations, which the README positions Mctx against directly: they are fast, but the README argues the speed "can come at the expense of usability and hackability." Mctx's claim is a middle position, a balance between performance and usability for researchers working in Python.

If you want to see the pattern applied, the README lists community projects: Pgx for vectorized JAX environments including an AlphaZero example, mctx_learning_demo for AlphaZero on random mazes, a0-jax for Connect Four, Gomoku and Go, muax for MuZero on CartPole and LunarLander, mctx-classic for a plain MCTS example on Connect Four, and mctx-az for AlphaZero with subtree persistence. Reading one of those before writing your own recurrent_fn is cheaper than reading the library source.

Maintenance, licence and the cost of upgrading

The repository is not archived, and the last push was on 2026-09-10, which is recent. Releases are infrequent and the version numbers are small: v0.0.71 on 2026-06-15, v0.0.6 on 2025-09-02, and v0.0.5 on 2023-11-24. The jump from 0.0.6 to 0.0.71 is unusual numbering, and the README does not explain it. What the release history does show is that you should not expect a steady cadence; plan to pin a version rather than track main.

The upgrade cost is concentrated in the JAX pin. pyproject.toml requires jax>=0.7.0 and jaxlib>=0.7.0. Any project that shares an environment with Mctx inherits that floor, and a future release could raise it again. Because the search is jitted, a JAX upgrade can also change compiled behaviour in ways that are not visible from Mctx's own changelog. Pinning both Mctx and JAX together is the safer arrangement.

On licensing, the repository ships an Apache-2.0 LICENSE file and pyproject.toml declares the Apache Software License classifier, so the code is permissively licensed. That is a statement about the project's terms, not advice about your situation; if you are embedding Mctx in a product, read the LICENSE file and the NOTICE requirements yourself, and note that the README's citation block asks for an academic citation when you publish work that uses it.

Editorial conclusion

Adopt Mctx if you already train in JAX and can express your environment as a recurrent_fn plus a RootFnOutput, and if you want search to run batched and jitted rather than in a Python loop. Do not adopt it if you need a turnkey opponent for a board game or you cannot supply learned priors and values, since Mctx has no environment of its own. Before committing, verify your JAX and jaxlib are at least 0.7.0 as pyproject.toml requires, confirm your environment can be written as a pure function of (params, rng_key, action, embedding), and run examples/policy_improvement_demo.py to see whether the policy improvement claim holds under your own value estimates.

Frequently asked questions

How do I install Mctx?

Install the released version from PyPI with pip install mctx, or the development version with pip install git+https://github.com/google-deepmind/mctx.git. It requires Python 3.11 or newer, plus jax and jaxlib at version 0.7.0 or above.

Does Mctx include environments or game rules?

No. Mctx provides the search only. You supply the root state as a RootFnOutput and the environment dynamics as a recurrent_fn that returns a RecurrentFnOutput and the next embedding.

Which policy should I use, muzero_policy or gumbel_muzero_policy?

The README recommends gumbel_muzero_policy, because Gumbel MuZero guarantees a policy improvement if the action values are correctly evaluated. The repository demonstrates this in examples/policy_improvement_demo.py.

Official sources

  1. google-deepmind/mctx on GitHub
  2. Issues
  3. License: Apache-2.0
  4. README
  5. Releases
Add this badge to your README

If you maintain this project, the badge below links readers to this analysis and shows its maintenance status from the daily GitHub snapshot. Paste the markdown into your README; add ?metric=license or ?metric=stars to the image URL for a different field.

Add this badge to your README

markdown
[![Hysen Labs](https://hysenlabs.com/badge/google-deepmind-mctx.svg)](https://hysenlabs.com/projects/google-deepmind-mctx)