Shape Rotation 101: An Intro to Einsum and Jax Transformers
sankalp.bearblog.dev
sankalp.bearblog.dev
The result is a couple of dense lines but one cannot just read them without going into a deep analysis for each line.
It is a pity that this has been accepted as the standard for machine learning. Worse, now every package has its own variant of NumPy (e.g. "import jax.numpy as jnp" in the article), which is incompatible with the standard one:
https://jax.readthedocs.io/en/latest/jax.numpy.html
I really would like a simpler array library that does stricter type checking, supports saner type specifications for composite types, does not broadcast automatically (except perhaps for matrix * scalar) and does one operation at a time. Casting should be explicit as well.
Bonus points if it isn't tied and inextricably linked to Python.
The kinds of type safety you want might be good for other use cases but for ML research they get in the way too much.
Also consider using "None" instead of "np.newaxis". To newcomers it's not as self-explanatory but it results in more readable code imho.
Otherwise - great article! Didn’t know this exists in numpy. A really neat way to express matrix operations.
(OTOH, I’m not an einsum expert; please feel free to delight me by pointing out how it’s possible to do these sorts of things :-)