3 ms·
So the naive MCTS implementation in Python is ridiculously inefficient. Of course, you could reimplement it in C++ but this then requires you to use the C wrapp
by brilee 4y ago
So the naive MCTS implementation in Python is ridiculously inefficient. Of course, you could reimplement it in C++ but this then requires you to use the C wrappers of Tensorflow/JAX to do the MCTS/neural network interop.
I came up with a nifty implementation in Python that outperforms the naive impl by 30x, allowing a pure python MCTS/NN interop implementation. See https://www.moderndescartes.com/essays/deep_dive_mcts/ https://www.moderndescartes.com/essays/deep_dive_mcts/
MCTX comes up with an even niftier implementation in JAX that runs the entire MCTS algorithm on the TPU. This is quite a feat because tree search is typically a heavily pointer based algorithm. It uses the object pool pattern described in https://gameprogrammingpatterns.com/object-pool.html https://gameprogrammingpatterns.com/object-pool.html to serialize all of the nodes of the search tree into one flat array (which is how it manages to fit into JAX formalisms). I suspect it's not a particularly efficient use of the TPU, but it does cut out all of the CPU-TPU round trip latency, which I'm sure more than compensates.
- cgreerrun 4y ago> I came up with a nifty implementation in Python that outperforms the naive impl by 30x, allowing a pure python MCTS/NN interop implementation. See https://www.moderndescartes.com/essays/deep_dive_mcts/ https://www.moderndescartes.com/essays/deep_dive_mcts/ Great post! Chasing pointers in the MCTS tree is definitely a slow approach. Although typically there are ~ 900 "considerations" per move for alphazero. I've found getting value/policy predictions from a neural network (or GBDT[1]) for the node expansions during those considerations is at least an order of magnitude slower than the MCTS tree-hopping logic. [1] https://github.com/cgreer/alpha-zero-boosted https://github.com/cgreer/alpha-zero-boosted