How to scale your model: A systems view of LLMs on TPUs
jax-ml.github.io
jax-ml.github.io
[0] Although PyTorch arguably encompasses 2 levels, with both a pure functional library like the JAX API, as well as a "neural network" framework on top of it. Whereas JAX doesn't have the latter and leaves that to separate libraries like Flax.
Are you suggesting that XLA would be where this "lower level" approach would reside since it can do more automatic optimization?
AST parsing via reflection means your ML compiler needs to re-implement all of Python, which is not a small language. This is a lot of work and hard to do well with abstractions that are not designed for those use-cases. (I believe Julia's whole language auto-diff systems struggle for essential the same reason.)
I literally am a paid ML compiler engineer and I have no idea what this means. You understand that reflection, ala looking in a mirror is about being about to identify a type's type at runtime. It has nothing to do with the AST.
https://docs.scala-lang.org/scala3/reference/metaprogramming...
Wikipedia: "reflection is the ability of a process to examine, introspect, and modify its own structure and behavior."
Would you say inspect.getsource(func) fits the definition of reflection?
Would you say ast.parse(inspect.getsource(func)) has something to do with the AST?
I would say that reflection is absolutely meaningless in an an interpreted runtime because you can always query the runtime.
> Would you say ast.parse(inspect.getsource(func)) has something to do with the AST?
It has something to do with the AST but it doesn't have much to do with reflection.
To what degree is this actually true, and what else is on the horizon that might become as popular as transformers?