google-deepmind / google-deepmind/mctx

`loop_fn` in `search.py` becomes slow as the tree depth increases

Open
#107 1 comment 0 reactions 0 assignees View on GitHub
Dominant language
Python
Stars
2.7k
Forks
218
Avg merge
23h 43m
Merged PRs (30d)
1

Description

I am using mctx to implement MuZero, where actions are selected at each node via a forward pass of a large neural network.
Initially, the main bottleneck in terms of time cost is the model forward pass, which is expected and similar to cases where C++-based MCTS is used.
However, after a few steps of training, as the model starts to learn certain biases, the search tree becomes deeper. At this point, the bottleneck shifts to the `loop_fn` used in the selection and backpropagation phases.

Here's the trace of one search. You can see that most of the time is spent in `search.py:291` and `search.py:181`, which correspond to (`search.py` was slightly adjusted to fit my codebase, so the line number mismatches):
- `tree, _, _ = jax.lax.while_loop(cond_fun, body_fun, loop_state)` in `backward`
- `end_state = jax.lax.while_loop(cond_fun, body_fun, initial_state)` in `simulate`

Image

Image

My question is: is there any way to improve the time efficiency of this part?
In practice, even with a tree depth of only a few tens, the performance is heavily affected by the huge gap in efficiency between `jax.lax.fori_loop` on GPU and CPU (with GPU being thousands of times slower).

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.