6 ms·
4000x Speedup in Reinforcement Learning with Jax
- _hark 4y agojax.vmap() is all you need?
- schizo89 4y agoNot only vectorization, but the plethora of environments written in jax. Hopefully someone will port MuJoCo to jax soon
- erwincoumans 4y agoThere is Brax, a differentiable physics simulator written in Jax. It includes Gym tasks such as Ant, Humanoid and more: https://github.com/google/brax https://github.com/google/brax It is not full MuJoCo but a good base to add more features. Aside from position based dynamics (xpbd) it features motion in generalized coordinates using the same accurate robot dynamics algorithms as MuJoCo and TDS (Tiny Differentiable Simulator).
- deleted 4y ago[deleted]
- schizo89 4y agoNeural differential equations are also easier with jax. sim2real may be easier with simulator where some of hard computations are replaced with neural approximations
- sillysaurusx 4y agoIt's a little disingenuous to say that the 4000x speedup is due to Jax. I'm a huge Jax fanboy (one of the biggest) but the speedup here is thanks to running the simulation environment on a GPU. But as much as I love Jax, it's still extraordinarily difficult to implement even simple environments purely on a GPU. My long-term ambition is to replicate OpenAI's Dota 2 reinforcement learning work, since it's one of the most impactful (or at least most entertaining) use of RL. It would be more or less impossible to translate the game logic into Jax, short of transpiling C++ to Jax somehow. Which isn't a bad idea – someone should make that. It should also be noted that there's a long history of RL being done on accelerators. AlphaZero's chess evaluations ran entirely on TPUs. Pytorch CUDA graphs also make it easier to implement this kind of thing nowadays, since (again, as much as I love Jax) some Pytorch constructs are simply easier to use than turning everything into a functional programming paradigm. All that said, you should really try out Jax. The fact that you can calculate gradients w.r.t. any arbitrary function is just amazing, and you have complete control over what's JIT'ed into a GPU graph and what's not. It's a wonderful feeling compared to using Pytorch's accursed .backwards() accumulation scheme. Can't wait for a framework that feels closer to pure arbitrary Python. Maybe AI can figure out how to do it.
- Inufu 4y agoAlphaZero did not run game logic on TPUs (neither chess nor other games), implementing it in C++ is more than fast enough and much simpler. TPUs were used for neural network inference and training, but game logic as well as MCTS was on the CPU using C++. JAX is awesome though, I use it for all my neural network stuff!
- sillysaurusx 4y agoAccording to the AlphaZero paper (https://arxiv.org/pdf/1712.01815.pdf https://arxiv.org/pdf/1712.01815.pdf) they ran game logic on TPUs: > Training proceeded for 700,000 steps (mini-batches of size 4,096) starting from randomly initialised parameters, using 5,000 first-generation TPUs to generate self-play games and 64 second-generation TPUs to train the neural networks. Further details of the training procedure are provided in the Methods.
- luchris429 4y agoAuthor here! I didn't realize this got posted on HN. While indeed we do get a speedup by putting the environments on the GPU, most of the speedup seems to come from the ability to easily parallelize RL training with Jax. While there is work on putting RL environments on accelerators, the main speedup from this work comes from also training many RL agents in parallel. This is largely because the neural networks we use in RL are relatively small and thus don't utilize the GPU very efficiently. While this was always possible to do, Jax makes it way easier because we just need to call `jax.vmap` to get it to work.
- percentcer 4y agoReminds me of this evergreen tweet from ryg: https://mobile.twitter.com/rygorous/status/1271296834439282690 https://mobile.twitter.com/rygorous/status/12712968344392826... if you made something 2x faster, you might have done something smart if you made something 100x faster, you definitely just stopped doing something stupid
- hyperbovine 4y agoMeh. This tweet is a lot less clever than it seems. Shave a factor of n off the complexity of your your algorithm, as happens regularly in CS and informatics, and have all the 1000x speedups you want.
- squeaky-clean 4y agoIf you shave a factor of n off of your algorithm, it usually isn't the same algorithm anymore. That's what they mean, the previous algorithm choice was "stupid" and you've stopped doing something stupid.
- whatshisface 4y agoWith that kind of creative interpretation almost any statement can be true. :)
- Ar-Curunir 4y agoThat's like saying, everybody struggling to solve SAT problems is just being stupid; just prove P = NP and solve the damn thing!
- squeaky-clean 4y agoIt's just an idiom and not meant to be taken so literally. Basically if you're getting a 100x improvement, you're jumping from one paradigm to another. And those are probably already solved algorithms / invented datastructures. If you're doing something novel and "clever" to speed up your algorithm for your circumstances, it's probably going to be under a 5x improvement. Or to continue the SAT analogy, if you're doubling your SAT score (e.g. 700 to 1400) you're getting smarter. If your 100x'ing your SAT score, the best possible previous score you could have gotten is a 16. Which would mean you got about one and a half questions correct across both the reading and math SAT.
- ssivark 4y agoFrom what I understand of Jax, it feels somewhat similar in flavor to Julia, but trying to live with the language constraints (and ecosystem benefits) of Python. I wonder how Julia is placed for running reinforcement learning algorithms (efficiently) — particularly in cases when the “environment” is nicely wrapped in Python to fit some standardized interface.
- Buttons840 4y agoI've done some RL experiments in Julia, and having all the in-between be fast was helpful; I saw significant speed increases. That said, Julia was probably just compensating for my own stupidity, because I was converting my environment objects into tensors over and over and over.
- ipsum2 4y agoStrange that the author claims Jax's vmap is what's doing the heavy lifting, but doesn't use PyTorch vmap to make the benchmark comparable.
- xyzzy4747 4y agoHow does this compare with PyTorch / Tensorflow / etc.? Obviously doing heavy data processing on the GPU will have a large speedup compared to a single thread on the CPU. It's almost like the author is claiming credit for creating Nvidia, when in fact he is just calling its APIs.
- m00x 4y agoI don't get that sense at all. You could do the same with tensorflow and pytorch, but in my experience, with more difficulty since they're more opinionated about how you should do your operations. JAX is definitely easier to do things that aren't on rails.
- luchris429 4y agoThe baseline we are comparing to is standard RL training that is widely used in academia. The technique mentioned in the blog post is not widely used amongst researchers. The reason we write about Jax is that doing this technique is really hard in PyTorch / Tensorflow. This is because: 1. Jax has vmap. (PyTorch does now too, but it is far more recent). 2. There are RL environments that others have written in pure Jax (see the blog post for four different repos of RL environments) 3. As m00x hints to, Jax replicates Numpy's API. This makes it way easier to use for non-neural network programming (e.g. RL environments).
- nothrowaways 4y agoIt is misleading, the speedup is not just because it is Jax. The devil is in the GPU
- luchris429 4y agoIndeed the devil is in the GPU! Jax and its ecosystem just make it much easier to use the GPU.