5 ms·
Has anyone benchmarked Jax? Curious how it compares to PyTorch for nontrivial networks, say ResNet.
by tbenst 7y ago
Has anyone benchmarked Jax? Curious how it compares to PyTorch for nontrivial networks, say ResNet.
- komuher 7y agoIt is compiled to XLA so should be a lot faster then pure PyTorch but probably will be slower then TVM (https://tvm.apache.org/ https://tvm.apache.org/) i can prepare some benchmarks in next few days if u are interested :)
- tbenst 7y agoThat would be wonderful if you’re able to! Also doubles as a good intro to Jax :). Please feel free to tweet at me (@tbenst) or email [same username at stanford dot edu] if you do get around to it.
- komuher 7y agoSure (dont have twitter yet) but will post it here on hacker news in next week probably :)
- p1esk 7y agoI'm very interested in seeing Resnet results. Especially FP16 precision running on tensor cores (V100 cards). Please use synthetic input vectors so that the input pipeline is not a bottleneck (as it is often the case, and it varies per framework).
- m0zg 7y ago> XLA so should be a lot faster I've yet to see anything get "a lot faster" because of XLA. It's a ton of complicated code, but then you end up spending the vast majority of time in NVIDIA's cuDNN anyway, so any benefits you might have hoped for will be marginal at best.
- p1esk 7y agohttps://blog.exxactcorp.com/nvidia-quadro-rtx-8000-deep-learning-performance-benchmarks-for-tensorflow-2019/ https://blog.exxactcorp.com/nvidia-quadro-rtx-8000-deep-lear... Almost double speedup for FP16 Resnet-50.
- m0zg 7y agoEasily outperformed by the more traditional TensorRT, which TF also supports: https://devblogs.nvidia.com/tensorrt-integration-speeds-tensorflow-inference/ https://devblogs.nvidia.com/tensorrt-integration-speeds-tens.... In fact, also seems to be outperformed by plain PyTorch using a single V100: https://github.com/NVIDIA/DeepLearningExamples/tree/master/PyTorch/Classification/ConvNets/resnet50v1.5#inference-performance-results https://github.com/NVIDIA/DeepLearningExamples/tree/master/P...
- p1esk 7y agoInteresting. I wonder why there's such a difference between Nvidia Pytorch benchmark and Exxact results: Nvidia is more than twice faster for single GPU. V100 should only be ~10% faster than Quadro 8000. Either Exxact is incompetent, or Nvidia has some special sauce.
- m0zg 7y agoFWIW, NVIDIA TensorRT pre-profiles the models before it runs them. I don't know how it does that exactly (that part is closed source) but I'd guess they just try different algorithms on each op individually (i.e. plain conv vs Winograd) and pick a good balance of speed and memory usage according to heuristics. On some nets this can make all the difference in the world, and ResNet50 is basically the most studied architecture in existence, so you can bet it's in every single benchmark for this kind of thing, and as such it receives disproportionate attention.
- p1esk 7y agoI thought all frameworks can do this type of profiling (e.g. torch.backends.cudnn.benchmark = True). Nvidia might have eliminated any potential data pipeline bottlenecks (with careful DALI tuning), but I'd still expect a lot less speedup. Maybe they compiled pytorch with certain tricks, and used newer CUDA/CuDNN code, idk.
- gdahl 7y agoCheck out the Flax ResNet50 example: https://github.com/google/flax/tree/master/examples/imagenet https://github.com/google/flax/tree/master/examples/imagenet It runs about as fast as any of the other popular machine learning frameworks, occasionally faster. Disclaimer: I work for Google and use JAX, although I'm not on the Jax team.