Live data from Hacker News

4000x Speedup in Reinforcement Learning with Jax

chrislu.page

1–10 of 32 posts

Re: 4000x Speedup in Reinforcement Learning with Jax

#5
It'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.

Re: 4000x Speedup in Reinforcement Learning with Jax

#7

Reminds me of this evergreen tweet from ryg: 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

Meh. 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.

Re: 4000x Speedup in Reinforcement Learning with Jax

#8

Reminds me of this evergreen tweet from ryg: 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

Meh. 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.

If 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.

Re: 4000x Speedup in Reinforcement Learning with Jax

#9

Reminds me of this evergreen tweet from ryg: 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

Meh. 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.

I think it’s fair to classify using O(n^2) sorting as stupid (if the input size is significant). In numerical computing, using a naive routine for computing eigenvalues in O(n^4) would equally be considered stupid, unless the input sizes are known to be small of course.

Re: 4000x Speedup in Reinforcement Learning with Jax

#10

Reminds me of this evergreen tweet from ryg: 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

Meh. 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.

Just prove P = NP while you are at it
Post reply on HN