Trade-Offs in Automatic Differentiation: TensorFlow, PyTorch, Jax, and Julia
stochasticlifestyle.com
stochasticlifestyle.com
Formally, there is a generalisation of differentiation which can handle functions like ReLU (i.e. locally Lipschitz non-differentiable functions) by allowing a derivative to be set-valued. It's called the Clarke gradient. The Clarke gradient of ReLU at 0 is the closed interval [0,1]. Note that the Clarke gradient doesn't satisfy the chain rule (except in a weakened form) which might seriously mess up some assumptions about autodiff. Is this generalised derivative useful in autodiff?
I imagine that this is a largely theoretical tool that's useful in analysing algorithms but useless for actually computing things.
[edit]
Question: Are there numerical applications in which the subgradient is actually computed, or is it a purely analytical tool?
See for example : https://www.stat.cmu.edu/~ryantibs/convexopt-F15/lectures/07...
and related neat usages in set based optimization methods in MathOptInterface (part of JuMP.jl): https://matbesancon.xyz/post/2020-12-24-chains_sets2/
The gist of it is that we endeavour to provide ‘useful’ values for the gradient at non-differentiable points, even if the traditional derivative is not defined or infinite.
You could imagine it as us smoothing out edges or discontinuities, ideally in a way that that makes things like gradient descent well behaved.
Wolfram Language takes an approach similar to Apple’s in providing a good number of pre trained models, but I haven’t yet discovered any automatic differentiation examples.
Of the frameworks described in the article, I find Julia most interesting but I need to use Python and TensorFlow in my work.
Following this thread, you can also see what how the Julia tools evolved. If you see the paper that was the synthesis for Zygote.jl, it was all about while loops and scalar operations (https://arxiv.org/abs/1810.07951). Why did that not completely change ML? Well, ML doesn't use those kinds of operations. I would say the project kind of started as a "tool looking for a problem". It did get a bit lucky that it found a problem: scientific applications need to be able to use automatic differentiation without rewriting the whole codebase to an ML library, leading to the big Julia AD manifesto of language-wide differentiable programming by directly acting on the Julia source itself rather than a language subset (http://ceur-ws.org/Vol-2587/article_8.pdf). Zygote was a good AD, but not a great AD, why? Because it could not hit this goal, mostly because of its lack of mutation handling. Yes, it does handle standard ML just fine, but does not justify its added complexity.
What has actually kept Julia AD research going is that some scientific machine learning applications, specifically physics-informed neural networks (PINNs), require very high order derivatives. For example, to solve the PDE u_t = u_xx with neural networks, you need to take the third derivative of the neural network. With Jax this can only be done with a separate language subset (https://openreview.net/pdf?id=SkxEF3FNPH), and thus a new AD for Julia to replace Zygote, known as Diffractor.jl, was devised to automatically incorporate higher order AD optimizations as part of the regular usage (https://www.youtube.com/watch?v=mQnSRfseu0c). It is these PINN SciML applications that have funded its development and is its built-in audience: it solves a problem nothing else does, even if it is potentially niche. Similarly with Enzyme, it solved the problem of how to do mutation well, which is where you can see in the paper that its applications are mostly ODE and PDE solvers (Euler, RK4, the Bruss semilinear PDE) (https://proceedings.neurips.cc/paper/2020/file/9332c513ef44b...). Torchscript and Jax do not handle this domain well, so it has an audience, which may (or may not?) be niche.
A big part of writing this blog post was to highlight this to the Julia AD crew that I regularly work with. What will keep these projects alive is understanding the engineering trade-offs that are made and who the audience is. The complexity has a cost so it better have a benefit. If that target is lost, if any benefit is a theoretical "but you may need more features some day", then the projects will lose traction. The project needs to be two-fold: identify new architectures and applications that would benefit from expanded language support from AD and build good support for those projects. Otherwise it is just training a transformer in Julia vs training a transformer in Python, and that is not justifiable.
Diffractor.jl has a much loftier goal: optimized differentiable programming of any code from any package in the Julia ecosystem. Because it's building typed IR, it will need a full set of Julia-based analysis tools (escape analysis, loop-invariant code motion, etc.) to approach the amount of optimization XLA can do when XLA optimizations are applicable. While such passes are being developed (for example, this is the PR for putting immutable array optimizations into the language so that Diffractor-friendly immutable can generate the optimized mutable form: https://github.com/JuliaLang/julia/pull/42465), it's at least a few years away before it's doing something like reliably combining multiple matrix-vector products into a matrix-matrix BLAS3 call. That would put it on even footing today to compete against PyTorch in a kernel vs kernel optimization battle, but not against TensorFlow code in cases where XLA optimizations are doing something more.
"Just wait 3 years and it will be really cool" is not a good way to start building a robust community, instead those interested in it need to ask how to demonstrate the improvements afforded by the added generality today. That's why what's making the project tick right now is the ARPA-E projects for physics-informed neural networks, the DJ4Earth project to do direct differentiation of the CLIMA climate model without changing any of the model code (https://dj4earth.github.io/), etc. Those kinds of projects are what is keeping a lot of the dev team open to be full time on these AD and compiler optimization projects. But if successful, it will also give an AD that is great for standard ML.
the advantage of transformers (computationally) seems to be how little sophistication the attention mechanism needs from AD systems (and how well it appears to scale with data). it's also a very static architecture in terms of a data flow/control flow perspective.
as far as I understand, this is far different from systems needing to be modeled in continuous time, especially things like SDEs. I am curious if things like delay embeddings will ever be modeled in terms of mechanisms similar to attention however.
Right now switching over to them would require a ton of code changes, relearning intuitions, debugging, profiling, etc. for not a ton of benefit.
sounds a lot like 'classical computer vision'. e.g, when I learned the subject (mid 2000s), topological features were all the rage: https://en.wikipedia.org/wiki/Digital_topology
Julia is designed for advanced numerical computing and Python isn't. The metaprogramming affordances needed for AD are much better developed in Julia than they ever will be in Python. And let's not forget the immense utility of multiple dispatch in Julia, another feature Python will probably never have. So it's not surprising that Julia is simply way more capable.
Python is just more approachable and natural to people. Julia should learn from that
But personally I find OOP ugly and unnatural, and Julia's model elegant and natural. And far more powerful - Julia programmers are using multiple dispatch to build out scientific computing to a sophistication not seen in any other language.
It might not be your cup of tea if you need to see object.method() in your code, but if you're more mentally flexible and want to build the next generation of technical computing tools, Julia is the place to be right now.
I’ve tried it for close to a year and the ergonomics still felt off, it reminds me of how the scala crowd talked about functional programming, and we’ve seen how that turned out.
I hear this from a lot of people that try Julia and yet the Julia crowds answer is always that they are dumb. Sounds a lot like the scala crowd…
Julia is far ahead in affordances to write fancy technical code and fairly behind in simple things, like standard affordances to write more ordinary code, or the ability to quickly load in data and make a plot.
I just think it's a misdiagnosis to blame multiple dispatch for this issue. It's much more about the Julia community prioritizing the needs of their target market.
The reason why Julia is fast is because automatic function specialization to concrete dispatches gives type-grounded functions which allows the complete optimization to occur on high-level looking code (see https://arxiv.org/abs/2109.01950 for type-theoretic proofs). It's basically a combination of (1) define a type system in a way that allows for type-grounded functions and compile-time shape inference (shape as in, byte structure of the structs), (2) define a multiple dispatch system with automatic function specialization on concrete types, (3) have a typed IR which proves and devirtualizes all dispatches before hitting the LLVM JIT. If you simply slap the LLVM JIT on random code, you will not get that performance. But now because multiple dispatch is fundamental to performance in the language, the rest of the "game" for the language is how to design an ergonomic language around this feature and how to teach people to use it effectively as a problem solving tool.
You actually see something similar going on in the world of Jax. With Jax, you need to be able to perform abstract interpretation to the Jax IR. In order for this to be possible with the interpreters Jax has, the functions that are being interpreted need to always have the same computational graph for the same inputs, i.e. they need to be pure functions. This is why Jax is built on functional programming paradigms. It would be similarly uncharitable to say the reason why Jax does not embrace OO is because the developers just love functional programming: the programming paradigm choice clearly falls out of what the tools needs to do.
It remains to be seen if Jax is the tool that makes more people finally embrace functional programming styles, or if enough people see pervasive performance necessary enough to change to the multiple dispatch paradigm of Julia. But what is clear is that tools that are moving away from OOP are not doing so arbitrarily, it's all about whether doing so is beneficial enough to justify the change.
This definitely fits with my experience. It took me quite a while to really "get" dispatch-oriented programming as a paradigm, but once I started to get it there was no going back.
But what actual problem is ORM solving beyond that?
Related to ORMs, but not quite on topic - query building. Type checked queries, parts of which can be passed around business logic, are very powerful and flexible.
That said, I find the concept of abstracting ML ingredients outside of languages a nice one although its not entirely novel(python's been doing this from day 1 :D). The strength for keeping it in 1 language can be profound though. Compilers can optimize across operations. Calling many atomic functions from an API/server from a client loses that unless implemented carefully. That one language benefit is a big part of what Julia has to offer.
I could see a value addition statement being made if the "whole market" solution included a lot of goodies. But every time I think of what that looks like - I think it looks like Julia in 2-5 years....
[1] - https://enzyme.mit.edu/
being able to do interprocedural cross language analysis seems awesome considering how much code is written in C++, but used in higher level languages.
There’s been some great work in this space in the past 5 years.
I’ve got some stuff I worked out this fall I’m overdue to write up and share some prototypes : there is a way to do reverse mode auto diff isolated to just being an invisible compiler pass! Without any of the extra complexity in what are otherwise equivalent formulations
[1]: https://github.com/breandan/kotlingrad/tree/c02ac55325c05a2e...
They called it differentiable programming https://ai.facebook.com/blog/paving-the-way-for-software-20-...
The Julia approach has been instead to expose an interface for compiler plugins that any AD (or other ‘nonstandard interpretation / code transformation) library can access. There’s negative trade offs to to this as well, but I really like the way it’s turned out and I think it’s given us some fantastic tools for non-AD purposes as well.
I assume you mean autograd?
For people who don't know, perturbation confusion is a problem that affects naive implementations of autodiff where derivatives of order bigger than 1 are not computed correctly.
With reverse mode AD, this is much less likely to be an issue because the AD system isn't necessarily storing and working on hidden extensions to the values, it's running a function forwards and then running a separate function backwards having remembered some values from the forward pass. If the remembered values are correct and never modified, then generating a higher order derivative is just as safe as the first. But that last little detail is thus what I think is most akin to perturbation confusion in reverse mode: reverse mode has the assumption that the objects captured in the forward pass will not be changed (or will be at least be back in the correct state) when it is trying to reverse. The easy way to break this assumption doesn't even require second derivatives. The easiest way to break it is mutation: if you walk forward by doing Ax, then the reverse pass wants to do A'v so it just keeps the pointer to A, but if A gets mutated in the meantime then using that pointer is incorrect. This is the reason why most AD systems simply disallow mutation except in very special unoptimized cases (PyTorch, Jax, Zygote, ...).
Enzyme.jl is an exception because it takes a global analysis of the program it's differentiating (with proper escape analysis etc. passes at the LLVM level) in order to know that any mutation going forward will be reversed during the reverse path, so by the time it gets back to A'*v it knows A will be the same. Higher level ADs could go cowboy YOLO style and just assume the reversed matrix is correct (and it might be a lot of the time), though that causes some pretty major concerns for correctness. The other option is to simply make a full copy of A every time you mutate an element, so have fun if you loop through your weight matrix. The Diffractor.jl near future approach is more like Haskell GHC where it just wants you to give it the non-mutating code so it can try and generate the mutating code when that would be more efficient (https://github.com/JuliaLang/julia/pull/42465).
So with forward-mode AD there was an entire literature around schemes of provable safety to perturbation confusion, and I'm surprised we haven't already started seeing papers about provable safety with respect to mutation in higher-level reverse-mode AD. I would suspect that the only reason why it hasn't started is that the people who write type-theoretic proofs tend to be the functional programming pure function folks that tell people to never mutate anyways, so the literature might instead go the direction of escape analysis proofs to optimize immutable array code to (and beyond) the performance of mutation code on commonly mutating applications. Either way it's getting there with the same purpose in mind.