PyTorch vs. TensorFlow in 2022
assemblyai.com
assemblyai.com
Just to make sure people aren't scared off by this: jax provides a lot more than just low level linear algebra. It has some fundamental NN functions in its lax submodule, and the numpy API itself goes way way past linear algebra. Numpy plus autodiff, plus automatic vectorization, plus automatic parallelization, plus some core NN functions, plus a bunch more stuff.
Jax plus optax (for common optimizers and easily making new optimizers) is plenty sufficient for a lot of NN needs. After that, the other libraries are really just useful for initialization and state management (which is still very useful; I use haiku myself).
I'm a pretty big fan of moving away from thinking about ML/Stats/etc specifically and people should more generally embrace the idea of differentiable programming as just a way to program and solve a range of problems.
JAX means that the average python programmer just needs to understand the basics of derivatives and their use (not how to compute them, just what they are and why they're useful) and suddenly has an amazing amount of power they can add to normal code.
The real power of JAX, for me at least, is that you can write the solution to your problem, whatever that problem may be, and use derivatives and gradient descent to find an answer. Sometimes this solution might be essentially a neural network, other times the generalized linear model, but sometimes it might not fit obviously into either of these paradigms.
Jax is definitely the right direction for the python ecosystem, but it can't solve all your problems. At some point you still need a fast language.
In particular i want to be able to measure feature importance on both inputs and internal layers on a sample by sample basis. This is the only thing currently holding me back from using JAX right now.
Alternatively a simle to read/understand/port implementation of DeepLIFT would work too.
thanks
Is that big industry lab Google or Deepmind? haha
Haiku is really cool - I haven't used Flax. It'll be really interested to see the development of JAX as time goes on. I also saw some benchmarks that show its neck-and-neck with PyTorch as the fastest of the three, but I think with more optimization its ceiling is higher than PyTorch's.
Of course. It's the only library that can be explained from first principles: https://jax.readthedocs.io/en/latest/autodidax.html
Even still, do you think researchers will want to take the time to learn all of that when PyTorch gives them no real reason to switch? Every day spent learning JAX is another day spent not reviewing literature, writing papers, or developing new models.
pytorch historically hasn't really focused on forward mode auto differentiation: https://github.com/pytorch/pytorch/issues/10223
this definitely limits its generality relative to jax, which makes it less than ideal for anything other than 'typical' deep neural networks
this is especially true when the research in question is related to things like physics or combining physical models and machine learning, which imho is very interesting. those are use cases that pytorch just isn't good at.
Every day spent learning JAX is also another day spent not trying to fit a round peg into a square hole of other libraries. I made the leap when I was doing things that were painful in pytorch. In terms of time, I think I came out ahead.
Not everything is a nail, and pytorch is better for some things, an jax is better for others. "Every day spent learning the screwdriver is a day spent not using your hammer."
To get started JAX is just knowing Python and adding `grad`, `jit` and `vmap` to the mix, it takes about 5 minutes to get going.
To me this is the real power of JAX, it can be viewed as a few functions that make it easy to take any python code you've written and work with derivatives using that. This gives it tremendous flexibility in helping you solve problems.
As an example, I mostly do statistical work with it, rather than NN focused work. It took probably a few minutes to implement a GLM with custom priors over all the parameters, and the use then Hessian for the Laplace approximation of parameter uncertainty. The proper way to solve this would have been using PyMC but this worked good enough for me, and building the model in scratch in JAX took less time than refreshing the PyMC api for me.
For example, a long chain of pmaps, each with some sort of device partitioning logic, not JIT compiling is extremely hard to understand. I basically had to binary search code until the compile errors disappeared.
I’m pretty sure tf is considered in maintenance mode within google as Brain and the tf creators themselves have moved to Jax. I do think Google learned a lot from tensorflow and am excited to see Jax pan out.
Pytorch is a pleasure to debug. I think pytorch jit could close the deployment gap.
TensorFlow seems to be spreading itself pretty thin. Maintaining so many language bindings, TensorFlow.js, TFlite, Server, etc. seem like they could all use some focus, BUT, and this is a big but, do you think if they can get each part of their ecosystem to an easily usable point that they'll have cornered the industry sector?
PyTorch is taking a much more targeted approach as seen with PyTorch Live, but I truly think that TFLite + Coral will be a game-changer for a lot of industries (and Google will make a fortune in the process). To me it seems like this is where Google's focus has lain in the AI space for the past couple of years.
What do you think?
I'd like to agree. Google was very far ahead of the curve when they released Coral. I was completely stoked when they finally added hardware video encoding to the platform with the release of the Dev Board Mini.
I want them to succeed but I fear if they don't drastically improve their Developer Experience, others will catch up and eat their lunch. TensorFlow has been hard to pick up. A few years ago when I was trying to pick this up to create some edge applications, PyTorch wasn't so much easier that it seemed worth sacrificing EdgeTPU support. But now PyTorch seems much, much easier than it did then, while TensorFlow hasn't seemed to improve in ease-of-use.
Now I'm genuinely considering sacrificing TFLite / EdgeTPU in favor of, say Jetson-esque solutions just so that I can start doing something.
Note: I am an amateur/hobbyist in this context, I am not doing Edge machine learning professionally.
Mostly this I suspect
Are you in research? I think TensorFlow's position in industry puts it in a kind of too-big-to-fail situation at this point. It'll be interesting to see what happens with JAX, but for now TensorFlow really is the option for industry.
Do you think TFLite + Coral devices will help breathe new life into TF?
Meanwhile PyTorch doesn't follow SemVer and always has breaking changes for every minor version increment. There's always "Backwards Incompatible Changes" section for every minor version release: https://github.com/pytorch/pytorch/releases
> tf.nn.conv2d(input, filters, strides, padding, data_format='NHWC', dilations=None, name=None)
> torch.nn.functional.conv2d(input, weight, bias=None, stride=1, padding=0, dilation=1, groups=1)
Have you checked out Google's Coral devices? DL has definitely been abused for marketing purposes, but I think the lack of delivery had more to do with the fact that DL was progressing far faster than the tools around them which make their intelligence actionable.
Part of this is because so many DL applications had to be delivered in a SaaS way, when local AI makes much more sense for a lot of applications. I think the TF -> TFLite -> Coral Device pipeline has the potential to revolutionize a LOT of industries.
I have done the tutorials and they all work. They seem to be very well maintained.
People I know said they never got the google Coral SDK working. Unfortunately they wouldn't give me their Corals. :(
Kind of offtopic but yeah, same. I'm a data scientist and right now I'm learning Django.
I don't get your point about JS developers not enjoying the fruits of these labors - they don't need to enjoy them because they work in a different domain. And if they're interested in playing around with deep learning, the higher level APIs are easy to pick up. I'm not sure what you're expecting to see.
I do feel like Google could do better communicating all of their different tools though. Their ecosystem is large and pretty confusing - they've got so many projects going on at once that it always seems like everyone gets fed up with them before they take a second pass and make them more friendly to newcomers.
Facebook seems to have taken a much more focused approach as you can see with PyTorch Live
Pretty cool imo
Even though I've been working with Tensorflow for a few years now and I feel like I do understand the API pretty well, to some extent that just means I'm _really_ good at navigating the documentation, because there's no way to intuit the way things work. And I still run into bizarre performance issues when profiling graphs pretty much all the time. Some ops are just inefficient - oh but it was fixed in 2.x.yy! Oh but then it broke again in 2.x.yy+1! Sigh.
However - and I know this is a bit of a tired trope, but any kind of industrial deployment is just vastly, vastly easier with Tensorflow. I'm currently working with ultra-low-latency model development targeting a Tensorflow-Lite inference engine (C-API, wrapped via Rust) and it's just incredibly easy. With some elbow grease and willingness to dive into low level TF-Lite optimisations, one can see end to end model inference times in the order of 10-100us for simple models (say, a fully connected dnn with a few million parameters), and between 100us-1ms for fairly complex models utilising contemporary architectures in computer vision or NLP. Memory overhead and control over inference computation semantics are easy.
As a nice cherry on top, we can take the same Tensorflow SavedModels that get compiled to TF-Lite files and instead compile them to tensorflow-js for easy web deployment, which is a great portability upside.
However, I know there's some incredible progress being made on what one might call 'environmental agnostic computational graph ILs' (on second thought, let's not keep that name) which should open up more options for inference engines and graph optimisations (operator fusion, rollups, hardware dependant stuff, etc).
Overall I feel like things have been continuously getting better for the last 5 years or so. I'm pleased to see so many more options.
Yes - a lot of TF users don't realize that knowing the "tricks of the trade" for wrangling TF don't apply in PT because it just works more easily.
I agree that industry-centric applications should probably use TF. TFX is just invaluable. Have you checked out Google's Coral devices? TFLite + Coral = revolution for a lot of industries.
Thanks for all your comments - I'm also really excited to see what the coming years bring. While we might debate if PT or TF is better, they're both undoubtedly improving very rapidly! So excited to see how ML/DL applications start permeating other industries
I basically don't believe you. I'm a researcher in this area (DNNs on FPGAs) and you cannot get these latencies on real models without going to FPGA (and you're not synthesizing Verilog from TF, unless you're one of my competitors...). Just your kernel launch overheads for GPU are on the order of 10ms. For example, here's a talk given at GTC a couple of years ago where they do get down to 35us (on tensorcores) using persistent kernels, but on a mickey mouse network
https://www.nvidia.com/en-us/on-demand/session/gtcsiliconval...
CPU (where you don't have to deal with async CUDA calls) won't save you either; again here's a paper from USENIX (so you know it's legit) that shows that lowest times for real networks on CPU are ~2ms (and that's on resnet18, far shy of "millions" of weights)
ML eng is my area of expertise, and I would advise strongly against tensorflow.
Both PyTorch and TensorFlow..."
Can an article really be any good if it starts off with such obvious SEO spam?
That's an interesting take. I fail to see how such mentions, in an article that compares two things and, thus, mentions both things together, are in any way SEO spam?
For example, in an article comparing apples and oranges I would expect to see a rather high number of mentions of "apples and oranges". After all, that is the topic.
PyTorch just recently took the lead. [0] So if I were having to choose between learning the either of them, I would go with PyTorch.
[0] https://trends.google.com/trends/explore?date=all&q=tensorfl...
So for me the choice is TF 2 because I can train models 5-10x faster using Google's TPUs than if I had used PyTorch. I know the PyTorch developers are working on TPU support but last I checked (this spring) it wasn't there yet and I wasn't able to make it work well on Google Colab.
JAX is totally new to me, is this Google's new Tensorflow in the future?
I help maintain https://github.com/capitalone/DataProfiler
Our sensitive data detection library is exported to iOS, android, and Java; in addition to Python. We also run distributed and federated use cases with custom layers. All of which are improved in tensorflow.
That said, I’d use pytorch if I could. Simply put, it has a better user experience.
The fact that PyTorch is pythonic and easier to debug makes it better for a ton of users, but TensorFlow keeps the entire DL process in mind more, not just modeling.
I asked a friend of mine at Google to sleuth around internally and get a sense for the health of the project. He said that it's used on some internal projects and seems to have a pretty healthy internal website. So hopefully it won't be cancelled soon.
Maybe. But I only get paid in whatever snacks I can scrounge from the employee fridge and pickings have been slim of late since so few people come into the office these days. I am down to the "expired mystery mozzarella cheese sticks" and the leftover ketchup packages from Woodranch BBQ & Grill that are in the drawer with all the unused chopsticks and plastic forks.
I too spoke to a friend at Google that is part of the team, and whilst he said there were no plans to cancel it, or make radical changes, when I asked about unplanned plans, he kinda just shrugged and said "You know Google..."
I have a dual solution approach, Mediapipe for "in use now" and OpenPose for validation, slower processing and the "Google just **ed us" moment we're both anticipating. I need to build my own pose analysis system, but right now I don't have the bandwidth.
On the last day of Christmas the CEO sent to me:
Thirty-two Manfrotto Tripod extenders
Sixteen Manfrotto tripods
Sixteen high speed cables
Sixteen Manfrotto C-clamps
Sixteen Manfrotto 3/8 to 1/4-20 reducers
Sixteen Quick release mounts
Sixteen 4K cameras
Fooouuuurrrrr high-speeeeeed PCIe capture cards
Three days to hit deadline
Two triggered circuit breakers
One really huge headache
And a new VR H.M.D.TF has more layer types, parallelizes better, is easier to assemble w keras, and you don't have to recreate the optimizer when loading from disk. pytorch doesn't have metrics out of the box. TF all the way.
github.com/aiqc/aiqc
Also, I think Lightning handles the issue of loading optimizers, but I'm not sure about that.
It's nice to see TF get some love, but I still think PyTorch has easier debugging and is more pythonic which lowers the barrier to entry for a lot of people.
Our pipeline is all PyTorch Lightning — this made development easy but we have been having numerous issues trying to leverage multiple GPUs (this is for sequence models), keep getting strange errors.
The stride>1 case has been a bit more controversial within TensorFlow, and there is ongoing discussion on the correct way to implement it within PyTorch on the issue: https://github.com/pytorch/pytorch/issues/3867
But keras is OK and I greatly appreciate that you can (sometimes) serialize everything to hdf5.
Do you think PyTorch can catch up here? I think Google's Coral devices give them a lock on embedded devices in the coming years
I’m not sure coral has enough of an edge to make it worthwhile relative to simpler edge deployment options like cpu
("Shakes old man's fist at sky, and at people who seems to enjoy boilerplate code too much")
I use PyTorch and TensorFlow, and the article is spot-on in regard to "mystery operations that take a long time" with no real rhyme or reason behind them with regard to TensorFlow. That said, on the whole, I skew more towards TensorFlow because it is generally easier to reason about the graph and how it connects. I also find the models that are available to usually be more refined, robust and useful straight out of the box.
With PyTorch I am usually fighting a slew of version incompatibilities in the API between even more point releases. The models often feel more slap-dash thrown together, research like projects or toy projects, and whilst the article points out the number of papers that use PyTorch far exceeds those that use TensorFlow, and the number of models for PyTorch dwarfs that of TensorFlow, there isn't a lot of quality in the quantity. "90% of everything is crap." Theodore Sturgeon. And that goes double for PyTorch models. A lot of the models, and even some datasets, just feel like throwaway projects that someone put up online.
If you are on macOS or Linux and using Python, PyTorch works fine, but don't step outside of that boundary. PyTorch and TensorFlow work with other operating systems, and other languages besides Python, but working with anything but Python when using PyTorch is a painful process fraught with pain. And yes, I expect someone to drop in and say "but what about this C++ framework?" or "I use language X with PyTorch every day and it works fine for me!" But again, the point stands, anything but Python with PyTorch is painful. The support of other languages for TensorFlow is far richer and far better.
And I will preface this with, "my knowledge may be out of date" but I've also noticed the type of models and project code available for TensorFlow and PyTorch diverge wildly once you get outside of the toy projects. If you are doing computer vision, especially with video and people, and you are not working on the most simplest of pose analysis, TensorFlow offers a lot more options of stuff straight out of the box. PyTorch has some good projects and models, but they are mostly of the Mickey Mouse hobby stuff, or an abstract research project that isn't very robust or immediately deployable.
I use TensorFlow in my day-to-day job. All that said, I like PyTorch for its quirkiness, its rapid prototyping, its popularity, and the fact that so many people are trying out a lot of different things, even if they don't work particularly well. I use PyTorch in almost all of my personal research projects.
I expect in the future for PyTorch to get more stable and more deployable and have better tools, if it can move slightly away from the "research tool" phase it is currently in. I expect Google to do the usual Google-Fuck-Up and completely change TF for the worse, break compatibility (TF1 to TF2) or just abandon the project entirely and move on to the next new shiny.