PyMC: Theano Is Dead, Long Live Theano
pymc-devs.medium.com
pymc-devs.medium.com
As always, we're happy to accept to contributions! If you're looking to get involved, now is a great time. Please don't hesitate to speak up or reach out, either on the Theano-PyMC GitHub repo (https://github.com/pymc-devs/Theano-PyMC) or some other way (my website's in my bio).
In that regard, I am curious about why Tensorflow didn't work out. I understand Tensorflow version 1 implements a declarative mode that I guess is in many ways similar to Theano's. I'm assuming v2 still supports that mode, on top of the new eager mode -- is that the case? If so, was there some aspect of its implementation that made it unsuitable for PyMC?
In order to improve the performance and usability of PyMC, I believe we need to automate things at the graph level, and Theano is by far the most suitable for this--between the two, at least.
You can find some work along these lines in the Symbolic PyMC project (https://github.com/pymc-devs/symbolic-pymc); it contains symbolic work done in both Theano and TensorFlow.
Will the JAX backend and integration with external JAX modules mean that we’ll see improvements to PyMC3’s variational inference module? That would really increase the versatility of PyMC3 for probabilistic modelling in Python.
NumPyro / JAX / PyTorch just seems like the most versatile offering out there right now
JAX
It was just a quick demonstration of how easily one can use Theano as a generalized graph "front-end", while also preserving its more unique and programmable symbolic optimization capabilities. JAX was one of a few "backends" I considered, and, due to the JAX Python library, it also looked like the most straightforward one to implement first.
on the contrary, i recently was looking for python libraries for some bayesian computations for some greenfield development and pymc3 was at the top of the list. with statistical libraries, i prioritize well-tested, large community, and extensive documentation. if others share the same priorities, there's a long future for pymc3.
I welcome and applaud the choice of JAX as it shows a lot of promise with autograd and flexible execution targets.
Because this sounds very numba-esque. I always thought Jax was just a math library, that was slightly more usable than Numpy. This article makes it seem that code written in Jax ends up being significantly faster than Numpy...almost close to C
That's NUMBA territory
Also, Tensorflow probability is moving to JAX.
Numba is great for CPU-only code, JAX would be my goto solution to add GPU on top of a numpy code.
Theano had a numpy backend so moving to JAX was very straightforward and enabled GPU computations which is now a must for fast deep learning.
However, it is a trade-off, the jitting time will increase your run time (you can expect at least 10 seconds of compiling time) and JAX does not let you do in-place operations (however, it has workarounds).
If you can afford the time needed to jit your code then I highly recommend JAX.