marin-community/haliax: Named Tensors for Legible Deep Learning in JAX ·
Though you don’t seem to be much for listening, it’s best to be careful. If you managed to catch hold of even just a piece of my name, you’d have all manner of power over me. — Patrick Rothfuss, The Name of the Wind Haliax is a JAX library for building neural networks with named tensors, in the tradition of Alexander Rush's Tensor Considered Harmful. Named tensors improve the legibility and compositionality of tensor programs by using named axes instead of positional indices as typically used in NumPy, PyTorch, etc. Despite the focus on legibility, Haliax is also fast, typically about as fast as "pure" JAX code. Haliax is also built to be scalable: it can support Fully-Sharded Data Parallelism (FSDP) and Tensor Parallelism with just a few lines of code. Haliax powers Levanter, our companion library for training large language models and other foundation models, with scale proven up to 70B parameters and up to TPU v4-2048. Here's a minimal attention module implementation in Haliax. For
Haliax Though you don’t seem to be much for listening, it’s best to be careful. If you managed to catch hold of even just a piece of my name, you’d have all manner of power over me. — Patrick Rothfuss, The Name of the Wind Haliax is a JAX library for building neural networks with named tensors, in the tradition of Alexander Rush's Tensor Considered Harmful . Named tensors improve the legibility and compositionality of tensor programs by using named axes instead of positional indices as typically used in NumPy, PyTorch, etc. Despite the focus on legibility, Haliax is also fast , typically about
Explore this link on the map →related reading
- Build a Transformer in JAX from scratch: how to write and train your own models | AI Summertheaisummer.com
- Why You Should (or Shouldn't) be Using Google's JAX in 2023assemblyai.com
- irhum.github.io - Tensor Parallelism with jax.pjitirhum.github.io
- Tensor-Transformer Variants are Surprisingly Performant — LessWronglesswrong.com
- GitHub - jax-ml/jax: Composable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more · GitHubgithub.com
- Tensor Considered Harmfulnlp.seas.harvard.edu
- Annotated Research Paper Implementations: Transformers, StyleGAN, Stable Diffusion, DDPM/DDIM, LayerNorm, Nucleus Sampling and morenn.labml.ai
- microgptkarpathy.github.io
- PyTorch internals : ezyang's blogblog.ezyang.com
- Why You Should (or Shouldn't) be Using Google's JAX in 2023assemblyai.com
- Using JAX to accelerate our research — Google DeepMinddeepmind.com
- Overleaf Examplearxiv.org