Using JAX, numpy, and optimization techniques to improve separable image filters
bartwronski.com
bartwronski.com
If so, is this the beginning of the end of Tensorflow? I know Tensorflow is still top for production, but it is certainly rapidly losing followings in the research field, and Pytorch and now starting to focus on deployment as they know this is their weakness.
And jax is made by the autograd people.
I still see people prefering Tensorflow over Pytorch because they have the feeling that it is more mature for production use. Meanwhile Jax has not converged on a recommended deep neural network framework (it has the low level pieces).
At the moment its a great building block that researcher should probably know.
Another goal is to make JAX a great system for playing with things like this!
There's also "cross-country optimization" (https://www-sop.inria.fr/tropics/slides/EdfCea05.pdf) for mixing some forward-mode into reverse-mode to improve memory efficiency. Analogously to jax.checkpoint, we've only experimented with exposing that manually (in jax.jarrett, named because of https://arxiv.org/abs/1810.08297), and even then only for a special case. There's a lot to learn about, experiment with, and build!
The article says:
> Optimization of arbitrary functions is generally a NP-hard problem (there are no solutions other than exploring every possible value, which is impossible in the case of continuous functions)
It is true that optimization of arbitrary functions, or even many interesting classes of functions, is NP-hard. However, the definition given of NP-hard is incorrect, and in fact, on modern hardware, existing SMT solvers such as Z3 can solve substantial instances of many interesting NP-hard optimization problems, precisely because they do not explore every possible value. Moreover, it is in general possible (but again NP-hard) to use interval arithmetic to rigorously optimize functions on continuous domains (which seems to be what is meant), as long as they are not too discontinuous; the answer you get is only an approximation of the true optimum, but you can calculate it to any desired precision.
One particularly interesting class of optimization problems — because they are not NP-hard — are continuous linear optimization problems, which can be solved in guaranteed polynomial time using interior-point methods or usually in polynomial time using the "simplex method". Contrary to what you'd think from the quote from the article, going from continuous to discrete makes the problem NP-hard again. There is also a note in Dercuano surveying the landscape of existing software and methods for solving linear optimization problems; there's a lot of very powerful stuff out there.
It turns out that you can efficiently solve an enormous range of practical optimization problems by introducing a small number of discrete variables into a linear optimization problem, thus gaining most of the performance benefit of using a linear optimizer. I don't know if there's a way to get a reasonable perceptual result in a case like this with a linear optimizer, though.
> This is where various auto-differentiation libraries can help us. Given some function, we can compute its gradient / derivative with regards to some variables completely automatically! This can be achieved either symbolically, or in some cases even numerically if closed-form gradient would be impossible to compute.
Automatic differentiation is a specific approach to differentiation which is an alternative to symbolic differentiation and the older kind of numerical differentiation. What JAX does is automatic differentiation.