The future of Deep Learning frameworks
neel04.github.io
neel04.github.io
Viva PyTorch! (Jax rocks too)
In another timeline AI would have made Lua popular.
The best part is it trampled TensorFlow which I personally find obtuse.
I wonder if it'd have been hated more than Python is - especially with the 1-based indexing...
Python has always gotten hate for being super super slow and having an ugly syntax (subjective ofc, but I happen to agree)
Jax looks like something completely different to me. Maybe I’m dumb and probably not the target audience, but it occurs to me that very few people are. When I read about using Jax, I find recommendations for a handful of other libraries that make it more useable. Which of those I choose to learn is not entirely obvious because they all seem to create a very fragmented ecosystem with code that isn’t portable.
I’m still not sure why I’d spend my time learning Jax, especially when it seems like most of the complaints from the author don’t really separate out training and inference, which don’t necessarily need to occur from the same framework.
1. PyTorch is going all-in on torch.compile -- Dynamo is the frontend, Inductor is the backend -- with a strong default Inductor codegen powered by OpenAI Triton (which now has CPU, NVIDIA GPU and AMD GPU backends). The author's view that PyTorch is building towards a multi-backend future isn't really where things are going. PyTorch supports extensibility of backends (including XLA), but there's disproportionate effort into the default path. torch.compile is 2 years old, XLA is 7 years old. Compilers take a few years to mature. torch.compile will get there (and we have reasonable measures that the compiler is on track to maturity).
2. PyTorch/XLA exists, mainly to drive a TPU backend for PyTorch, as Google gives no other real way to access the TPU. It's not great to try shoe-in XLA as a backend into PyTorch -- as XLA fundamentally doesn't have the flexibility that PyTorch supports by default (especially dynamic shapes). PyTorch on TPUs is unlikely to ever have the experience of JAX on TPUs, almost by definition.
3. JAX was developed at Google, not at Deepmind.
1. I'm well aware of the PyTorch stack, but this point:
> PyTorch is building towards a multi-backend future isn't really where things are going
>PyTorch supports extensibility of backends (including XLA)
Is my problem. Those backends just never integrate well as I mentioned in the blogpost. I'm not sure if you've ever gone into the weeds, but there are so many (often undocumented) sharp edges when using different backends that they never really work well. For example, how bad Torch:XLA is and the nightmare inducing bugs & errors with it.
> torch.compile is 2 years old, XLA is 7 years old. Compilers take a few years to mature
That was one of my major points - I don't think leaning on torch.compile is the best idea. A compiler would inherently place restrictions that you have to work-around.
This is not dynamic, nor flexible - and it flies in the face of torch's core philosophies just so they can offer more performance to the big labs using PyTorch. For various reasons, I dislike pandering to the rich guy instead of being an independent, open-source entity.
2. Torch/XLA is indeed primarily meant for TPUs - like the quoted announcement, where they declare to be ditching TF:XLA in favour of OpenXLA. But there's still a very real effort to get it working on GPUs - infact, a lab on twitter declared that they're using Torch/XLA on GPUs and will soon™ release details.
XLA's GPU support is great, its compatible across different hardware, its optimized and mature. In short, its a great alternative to the often buggy torch.compile stack - if you fix the torch integration.
So I won't be surprised if in the long-term they lean on XLA. Whether that's a good direction or not is upto the devs to decide unfortunately - not the community.
3. Thank you for pointing that out. I'm not sure about the history of JAX (maybe might make for a good blogpost for JAX devs to write someday), but it seems that it was indeed developed at Google research, though also heavily supported + maintained by DeepMind.
Appreciate you giving the time to comment here though :)
> > torch.compile is 2 years old, XLA is 7 years old. Compilers take a few years to mature
> That was one of my major points - I don't think leaning on torch.compile is the best idea. A compiler would inherently place restrictions that you have to work-around.
There are plenty of compilers that place restrictions that you barely notice. gcc, clang, nvcc -- they're fairly flexible, and "dynamic". Adding constraints doesn't mean you have to give up on important flexibility.
> This is not dynamic, nor flexible - and it flies in the face of torch's core philosophies just so they can offer more performance to the big labs using PyTorch. For various reasons, I dislike pandering to the rich guy instead of being an independent, open-source entity.
I think this is an assumption you've made largely without evidence. I'm not entirely sure what your point is. The way torch.compile is measured for success publicly (even in the announcement blogpost and Conference Keynote, link https://pytorch.org/get-started/pytorch-2.0/ ) is by measuring on a bunch of popular PyTorch-based github repos in the wild + popular HuggingFace models + the TIMM vision benchmark. They're curated here https://github.com/pytorch/benchmark . Your claim that its to mainly favor large labs is pretty puzzling.
torch.compile is both dynamic and flexible because: 1. it supports dynamic shapes, 2. it allows incremental compilation (you dont need to compile the parts that you wish to keep in uncompilable python -- probably using random arbitrary python packages, etc.). there is a trade-off between dynamic, flexible and performance, i.e. more dynamic and flexible means we don't have enough information to extract better performance, but that's an acceptable trade-off when you need the flexibility to express your ideas more than you need the speed.
> XLA's GPU support is great, its compatible across different hardware, its optimized and mature. In short, its a great alternative to the often buggy torch.compile stack - if you fix the torch integration.
If you are an XLA maximalist, that's fine. I am not. There isn't evidence to prove out either opinions. PyTorch will never be nicely compatible with XLA until XLA has significant constraints that are incompatible with PyTorch's User Experience model. The PyTorch devs have given clear written-down feedback to the XLA project on what it takes for XLA+PyTorch to get better, and its been a few years and the XLA project prioritizes other things.
In the context of scientific computing - this is completely, blatantly false. We're not lowering low-level IR to machine code. We want to perform certain mathematical processes often distributed on a large number of nodes. There's a difference between ensuring optimization (i.e no I/O bottlenecks, adequate synchronization between processes, overlapping computation with comms) vs. simply transforming a program to a different representation.
This is classic [false analogy](https://simple.wikipedia.org/wiki/False_analogy)
Adding constraints does mean that you give up on flexibility precisely because you have to work around them. For example, XLA is constrained intentionally against dynamic-loops because you lose a lot of performance and suffer a huge overhead. So the API forces you to think about it statically (like you can work around it with fancier methods like using checkpointing and leveraging a tree-verse algorithm)
I'll need more clarification regarding this point, because I don't know what dev in which universe will not regard "constraints" as flying against the face of flexibility.
> popular HuggingFace models + the TIMM vision benchmark
Ah yes, benchmark it on models that are entirely static LLMs or convnet-hybrids. Clearly, high requirement on dynamicness and flexibility there.
(I'm sorry but that statement alone has lost you any credibility for me.)
> Your claim that its to mainly favor large labs is pretty puzzling.
Because large labs often play with the safest models, which often involves scaling them up (OAI, FAIR, GDM etc.) and those tend to be self-attention/transformer like workloads. The devs have been pretty transparent about this - you can DM them if you want - but their entire stack is optimized for these usecases.
And ofcourse, that won't involve considering for research workloads which tend to be highly non-standard, dynamic and rather complex and much, much harder to optimize for.
This is where the "favouring big labs" comes from.
> 1. it supports dynamic shapes
I agree that in the specifically narrow respect of dynamic shapes, it's better than XLA.
But then it also misses a lot of the optimization features XLA has such as its new cost model and Latency Hiding Scheduler (LHS) stack which is far better at async overlapping of comms, computations and even IO (as its lazy).
> there is a trade-off between dynamic, flexible and performance
Exactly. Similarly, there's a difference in the features offered by each particular compiler. Torch's compiler's strengths may be XLA's weakness, and vice-versa.
But its not perfect - no software can be, and compilers certainly aren't exceptions. My issue is that the compiler is being considered at all in torch.
There are use-cases where the torch.compile stack fails completely (not sure how much you hang around more research-oriented forums) wherein there are some features that simply do not work with torch.compile. I cited FSDP as the more egregious one because its so common in everyone's workflow.
That's the problem. Torch is optimizing their compiler stack for certain workloads, with a lot of new features relying on them (look at newly proposed DTensor API for example).
If I'm a researcher with a non-standard workload, I should be able to enjoy those new features without relying on the compiler - because otherwise, it'd be painful for me to fix/restrict my code for that stack.
In short, I'm being bottlenecked by the compiler's capabilities preventing me to fully utilize all features. This is what I don't like. This is why torch should never be leaning at a compiler at all.
It 'looks' like a mere tradeoff, but reality is just not as simple as that.
> XLA:GPU
I don't particularly care if torch uses whatever compiler stack the devs choose - that's beside the point. Really, I just don't like the compiler-integrated approach at all. The choice of the specific stack doesn't matter.
Jax's advantages shine when it comes to parallelizing a new architecture across multiple GPU/TPUs, which it makes much easier than PyTorch (no need for custom cuda/networking code). Needing to scale up a new architecture across many GPUs is however not a common use-case, and most teams that have the resources for large-scale multi-gpu training also have the resources for specialised engineers to do it in PyTorch.
The issue was TF had too many interfaces to accomplish the same thing and each one was rough in its own way. Along with some complexity for using serving and experiment logging via Tensorboard, but this wasn’t as bad at least for me.
Keras was integrated in an attempt to help, but ultimately it wasn’t enough and people started using Torch more and more even against the perception that TF was for prod workloads and Torch was for research.
TFA mentions the interface complexity as starting to be a problem with Torch, but I don’t think we’re anywhere near the critical point that would cause people to abandon it in favor of JAX.
Additionally with JAX you’re just shoving the portability problems mentioned down to XLA which brings its own issues and gotchas even if it hides the immediate reality of said problems from the end user.
I think the Torch maintainers should watch not to repeat the mistakes of TF, but I think theres a long way to go before JAX is a serious contender. It’s been years and JAX has stayed in relatively small usage.
If the future is going to be better more intelligent compilers, then that settles the question in my opinion.
Interesting take - I agree here somewhat.
But also, wouldn't you think a framework that has been from the ground-up designed around a specific, mature compiler stack be better able to integrate compilers in a more stable fashion than just shoe-horning static compilers into a very dynamic framework? ;)
Because JAX is not designed around a mature compiler stack. The history of Jax is more so that it matured alongside the compiler...
It is a messy and quickly expanding codebase with many surprises like segfaults and leaks.
Is scientific experimentation really sped up by these frameworks? Everyone uses the Transformer model and uses the same algorithms over and over again.
If researchers wrote directly in C or Fortran, perhaps they'd get new ideas. The core inference (see Karparthy's llama.c) is ridiculously small. Core training does not seem much larger either.
... then they would get nothing done.
"PyTorch is dead. Long live JAX." conveys exactly what the article about, and is a much better title.
Julia also competes in this domain from a more practical standpoint and has less limitations than JAX as I understand it, but is less mature and still working on getting wider traction.
Shameless plug for one of my talks at JuliaCon 2024: https://www.youtube.com/live/ZKt0tiG5ajw?t=19747s. The comparison between Python and Julia starts at 5:31:44.
Where do you feel Julia is at this point in time (compared to say, JAX or PyTorch) from a practitioner's standpoint?
I'm positive about Julia's future because the developer experience just feels so fun and productive. I always find it impressive how much a small group of self-organized volunteers has been able to achieve. Amazing things could happen if a company like Google or Meta paid a team of full-time engineers to advance the deep learning ecosystem. Fun fact: Julia strongly influenced PyTorch's recent design decisions [1].
[1]: https://dev-discuss.pytorch.org/t/where-we-are-headed-and-wh...
So not quite “JAX without limitations” — but certainly without some of the limitations.
Now it's been forever since I used PyTorch or TF, but I only remember TF 1.x being more like "why TF isn't this working." At some point I didn't blame myself, I blamed the tooling, which TF2 later admitted. It seemed like no matter how skilled I got with TF1, it'd always take much longer than developing with PyTorch, so I switched early.
Typescript and Elm would like a word
Also, generally when people complain that JS won the web, it's not because they prefer TS, it's cause they wanted to use something else and can't.
Never used Elm, but... no variables, kinda like Erlang, which I've used. That has its appeal, but you're not going to find a consensus that this is better for web.
I want to write the type, and for that to reveal the mistake.
Other than that, I haven't noticed their inference capabilities being any different.
So I don't get your point.
Can you give a single example where Rust has more automatic type inference compared with TS? (honest question, maybe I'm missing something)
That might be what you meant, but its not what you said. Thanks for clarifying.
PyTorch has better adoption / network effects. JAX has stronger underlying abstractions.
I use both. I like both :)
I think the biggest, well "con" I've seen is non-technical - the fear of JAX being killed by Google.
I mention in the blog as well [here](https://neel04.github.io/my-website/blog/pytorch_rant/#gover...) how important having an independent governance structure is. I'm sure for many big companies and labs, the lack of a promise of long-term, stable support is a huge dealbreaker.
I'm not sure how much Google bureaucracy would limit this, but have you raised the subject of forming an independent entity to govern JAX, very much like PyTorch? I believe XLA is protected, as its with TF governance. But perhaps, there could be one for JAX's ecosystem as well, encompassing optax, equinox, flax etc.
What concerned me about JAX, at a small company, is that it doesn't benefit from the network effects of almost everyone developing for it. E.g. There is no Llama 3.1 implementation in JAX afaict.
So as long as there is a need to pull from the rest of the world the ecosystem will trump the framework.
Activity in the LLM space is slowing down though, so there is an opportunity to take the small set of what worked and port it to JAX and show people how good that world is.
ive never even heard of jax nor will i have the skills to use it
i literally just want to know two things: 1) how much vram 2) how to run it on pytorch
its like fretting about how everything should be written in C++ instead of Python/Javascript
I don't care.
if not im not interested. i'll keep using pytorch.
i hope this answer makes more sense for you.
Downvoted. Hmmm. I’m a little tired so I don’t want to go into detail. However, I was a Perl programmer when Python was rising. So, needless to say, having a big lead doesn’t matter.
Please learn from history. A big lead means nothing.
So a good lesson is not to get distracted by foolish endevours.
KV caching is directly in conflict with a purely functional approach.
It is basically the same as memorization.
I'm glad they've removed the (rather arbitrary, and admittedly stupid) loc cap. And from the little I know, geohot is focusing on having its own internal compiler stack.
As much as I admire geohot, I don't think rolling your own compiler is the best way. Its not that the TinyGrad team isn't smart enough, but a compiler is a huge undertaking and that you have to support and maintain for a long time. I'm sure he's well aware of this, but no big labs would touch TG seriously because of this limitation.
XLA on the other hand is under governance seperate from Google, and is far more mature - so people trust that.
That said, I don't know much about Tinygrad so I would appreciate if someone more knowledgable can jump in here and outline the differences and key features ¯\_(ツ)_/¯
... from the company that pioneered the approach with tensorflow. I've worked with worse ML frameworks, but they're by now pretty obscure; i cannot remember (and i am very happy about it) the last time i saw MXNet in the wild, for example. You'll still find Caffe on some embedded systems, but you can mostly sidestep it.
Some assert-ing won't hurt you. Seriously. It might even help keeping your sanity.
Leaving aside the fact that PyTorch's ecosystem is 10x to 100x larger, depending on how one measures it, PyTorch's biggest advantage, in my experience, is that it can be picked up quickly by developers who are new to it. Jax, despite its superiority, or maybe because of it, can not be picked up quickly.
Equinox does a great job of making Jax accessible, but Jax's functional approach is in practice more difficult to learn than PyTorch's object-oriented one.
The subtitle doesn't convey the content of the article nearly as well as the title does. Perhaps you can take a sentence like "PyTorch has been a net negative for scientific computing efforts" which the article does say, or some toned down versino of the original title, but the current title makes it sound like a very different article and felt like clickbait to me.
Finally! Someone says it! This is why the C programming language will never have wide adoption. /s