Introduction to Haliax - Colab
Haliax is a JAX library for building neural networks with named tensors, in the tradition of Tensor Considered Harmful. We use named tensors in Levanter to improve the legibility and compositionality of our programs without sacrificing their performance or scalability. We're going to build a simple Transformer for an autoregressive language model, starting from basic Haliax concepts. We'll assume you have some familiarity with Transformers. This tutorial mainly focuses on the legibility aspects of named tensors as implemented in Haliax. We'll cover scaling, including fully-sharded data parallelism in the next tutorial. Haliax is available on PyPI with nightly dev builds. The typical way people build neural networks is with an neural net library like PyTorch, Tensorflow, Keras, or, in the Jax world, Flax, Haiku, or Equinox. All of these libraries are mainly centered around organizing n-dimensional arrays into modules, and then doing compute on the parameters and input data. I get really