PyTorch vs. TensorFlow in Academic Papers
horace.io
horace.io
However .. one area where Tensorflow shined was the static graph. As our models get even more intensive and needs different parts to execute in parallel, we are seeing some challenges in PyTorch's execution model. For example:
https://pytorch.org/docs/stable/notes/cuda.html#use-nn-paral...
It appears to me that high performance model execution is a bit tricky if you want to do lots of things in parallels. TorchServe also seems quite simple compared to offerings from Tensorflow. So in summary, I think Tensorflow still has some features unmatched by others. It really depends on what you are doing.
See https://blog.google/technology/ai/introducing-pathways-next-...
and Jeff Dean's TED talk: https://www.ted.com/talks/jeff_dean_ai_isn_t_as_smart_as_you...
What continues to surprise me is that for all his cleverness, Jeff Dean and the rest of the TF leadership spend the last 10 years basically recreating MPI-style high performance computing, but threw away all the learning and rebuilt every bit (except the matrix libraries) from scratch.
TF started with parameter servers (every machine has its own copy of weights and periodically contributes them to a common model, asynchronously) to models that are sharded by data input and model structure that is mapped to the TPU topology (TPUv4 is a 3D wrapped torus). Really not that different from the T3E I used in the 90s.
Word got around that debugging PyTorch was relatively painless, those earlier models made it into publications, and now here we are.
The #1 problem with PyTorch is that it’s great if you want to use one videocard for training. Facebook has completely failed to support research scientists that want to do more than this.
It’s no secret that I’m a jax fanboy. But I drink the koolaid because it tastes better than anyone else’s. PyTorch is gonna have a rude wake up call in about… oh, four years. They’ll wake up and hear everyone else comparing them to tensorflow, and it won’t be for the rosy reasons they currently enjoy. PyTorch devs are living in the dark ages without even realizing how much better it is when you have actual control over which parts of your program are JITed, along with an actual execution graph that you can walk and macroexpand lisp-style.
https://jax.readthedocs.io/en/latest/autodidax.html should be required reading for every ML dev, and I can hardly get anyone to look at it. Sometimes I wonder if people just don’t see the steamroller coming for PyTorch. Probably — jax still reads to outsiders as a toy.
Incorrect information so confidently stated here. Tons of research papers that use more than one GPU for training, not sure what you're referring to? Standard DDP works fine, for starters.
Can you elaborate? What’s the advantage of controlling which parts are jitted?
There is TorchServe but I haven't used it so I'm not sure how production ready it is. You have Nvidia's triton server which support cpu and gpu with tf1,tf2,pytorch,onnx and tensorRT.
You have onnx runtime which can run on cpu and gpu and there are convertors from tf and pytorch to onnx.
Then you have cloud based solutions like AWS sagemaker, elastic inference endpoints and even Inf1 instances that use AWS Inferentia chips which you would run with the Neuron SDK, they even have TensorFlow serving containers with built it support for Inferentia.
End of the day it really depends on your model, size, latency, inference runtime and the cost obviously.
And that's before optimizations like FP16, BFLOAT16, TF32, INT8, pruning, layers rewrite, getting rid of batch normalization etc.
Then you have up and coming solutions like Neural Magic (not associated) deepsparse to create sparse models for inference.
And that's just for cloud if you are talking about edge ml it's even more down the rabbit hole..
I know both Keras and PyTorch, and I will recommend PyTorch any day.
I haven’t used TF/Keras in the last 2.5 years outside of Edge AI projects- neither at work nor at personal projects.
JAX is interesting, though. But definitely not advisable as someone's first DL framework.
I am using Jax for differentiable programming, and in many cases, I saw enormous speedup after jit, sometimes in the ballpark of 1e4.
For Neural Networks, I use Equinox, and/or Elegy.
I'm curious what your workloads are that you're seeing speedups of as much as 1e4? Greatest I've heard of before was ~1e2 on some differential equation solving.
The 1e4 speedup was on a Trust Region optimizer. The algorithm was implemented to solve "the hard case"[1] involves multiple Cholesky factorizations, a matrix inversion, an eigenvalue decomposition on each step, and a call to scipy.linalg.solve_triangular.
Part of the speedup is likely from caching/avoiding recomputations things.
Granted, I had to rewrite a lot of the code to accommodate jax's peculiarities around python semantics, and made extensive use of jax.lax.{fori_loop, while_loop, scan, cond}.
[1] http://www.apmath.spbu.ru/cnsa/pdf/monograf/Numerical_Optimi..., see page 87.
Keras wasn't hard enough.
Pytorch was juuuuust right.