Useful algorithms that are not optimized by Jax, PyTorch, or TensorFlow
stochasticlifestyle.com
stochasticlifestyle.com
Tangentially, thinking about Julia, while one initially gets awed by the speed, and then the multiple dispatch, I wonder whether it’s deepest superpower (that we’re still discovering) might be the expressiveness to augment the compiler to do interesting things with a piece of code. Generic programming then acts as a lever to use these improvements for a variety of use cases, and the speed is merely the icing on the cake!
The framework of abstract interpretation, when combined with multiple dispatch as a language design feature, is absolutely insane.
I think programming language enthusiasts might meditate on these points —- and get quite excited with the direction that the Julia compiler implementation is heading.
Otherwise, if you’re curious about how Julia’s type inference algorithm works via abstract interpretation, I would recommend the following blog posts:
https://juliacomputing.com/blog/2016/04/inference-convergenc... https://juliacomputing.com/blog/2017/05/inference-converage2...
Julia has very interesting propositions on the subject, from language-level autodiff (https://fluxml.ai/Zygote.jl/latest/) to automated probabilistic programming (https://turing.ml/dev/) through DEs (https://diffeq.sciml.ai/stable/) and optimization (https://jump.dev/).
The whole ecosystem is in ebullition, and I'm very eager to see if it will be able to transform in the comping years into a solid foundation able to rival the layers of warts stacked on top of Python.
Zygote (for example) is based on a language level feature — generated functions which allow multistage programming in Julia — but beyond using this feature it does not configure or otherwise modify the normal compiler pipeline.
With Julia, it is sometimes tough to tell — but Diffractor (for example) I would consider a language level library, as it modifies the traditional compiler pipeline to perform inference and optimization in a specific way.
This discovered sparsity can be exploited by a lot the generic algorithms in Julia that also are able to efficiently work on sparse matrices. There is also a way of combining this with AD to get automatic sparse hessian of optimization which is huge.
I remember spending a summer using Template Model Builder (TMB), which is a useful R/C++ automatic differentiation (AD) framework, for working with accelerated failure time models. For these models, the survival to time T given covariates X is defined by S(t|X) = P(T>t|X) = S_0(t exp(-beta^T X)) for baseline survival S_0(t). I wanted to use splines for the baseline survival and then use AD for gradients and random effects. Unfortunately, after implementing the splines in template C++, I found a web page entitled "Things you should NOT do in TMB" (https://github.com/kaskr/adcomp/wiki/Things-you-should-NOT-d...) - which included using if statements that are based on coefficients. In this case, the splines for S_0 depend on beta, which is this specific excluded case:(. An older framework (ADMB) did not have this constraint, but dissemination of code was more difficult. Finally, PyTorch did not have an implementation of B-splines or an implementation for Laplace's approximation. Returning to my opening comment, there is no free lunch.
A longer answer: the splines require a basis matrix B(t), do that g(S_0(t))= B(t) gamma for a vector of parameters gamma and some transformation g of survival. A classical choice would be to use M-splines and I-splines with g(S)=-log(S) and a penalised likelihood, with the constraint that the gammas should be increasing. In R, this would use the splines2 package, while in Julia, one could use the splines2.jl package (disclaimer: which I maintain). The computational challenge is that the basis matrices need to be re-evaluated for changes in the coefficients for the covariates (that is, the betas).
It seems like you can't solve this kind of thing with a new jax primitive for the algorithm, but what prevents new function transformations from doing what the mentioned julia libraries do? It seems like between new function transformations and new primitives, you out to be able to do just about anything. Is XLA the issue, and you could run but not jit the result?
To get the more flexible form, you really would want to do it in a way that uses a full programming language's IR as its target. I think trying to use a fully dynamic programming language IR directly (Python, R, etc.) directly would be pretty insane because it would be hard to enforce rules and get performance. So some language that has a front end over an optimizing compiler (LLVM) would probably make the most sense. Zygote and Diffractor uses Julia's IR, but there are other ways to do this as well. Enzyme (https://github.com/wsmoses/Enzyme.jl) uses the LLVM IR directly for doing source-to-source translations. Using some dialect of LLVM (provided by MLIR) might be an interesting place to write a more ML-focused flexible AD system. Swift for Tensorflow used the Swift IR. This mindset starts to show why those tools were chosen.
And XLA has dynamic shape semantics (currently unused by jax) via SetDimensionSize: https://www.tensorflow.org/xla/operation_semantics#setdimens...
This seems to be the key bit. It’s a great data point around the meme of “with a sufficiently advanced compiler…” In this case we have sufficiently advanced compilers to make very different JIT trade offs. XLA is differently powerful compared to Julia. Very cool, thanks for the insight.