Build a Transformer with JAX

Google for DevelopersAbout 8 min readMay 22, 2025Watch original
THE SUMMARYAI-generated

Key Concepts:

  • General Purpose Transformer (GPT): A type of AI model that has significantly impacted the AI field.
  • Embeddings: Numerical representations of words or concepts that allow computers to process and interpret language.
  • Attention Mechanism: A core component of transformers that calculates the relationships between words in a context.
  • Query, Key, Value (Q, K, V): Vectors used in the attention mechanism to determine relationships between words.
  • Multi-Headed Attention: Multiple parallel implementations of the attention mechanism.
  • Skip Connections: Connections that bypass certain blocks in the model, allowing data from earlier layers to be combined with later layers.
  • JAX: A numerical computation library for high-performance machine learning.
  • Flax NNX: A neural network library built on top of JAX.
  • Optax: A gradient processing and optimization library for JAX.
  • Orbax: A library for saving and loading JAX model checkpoints.
  • XLA: A compiler for JAX that optimizes performance on accelerated hardware.
  • Data Parallelism: A strategy where the data is split across multiple devices, while the model is replicated on each device.
  • Device Mesh: A grid of devices used to distribute data and models for parallel processing.
  • Hyperparameters: Parameters that control the training process, such as batch size and learning rate.
  • Softmax: A function that normalizes values to a probability distribution.
  • Layer Normalization: A technique to normalize the inputs to a layer, improving training stability.
  • GELU: A type of activation function used in transformer models.
  • Weights & Biases (W&B): A platform for tracking and visualizing machine learning experiments.
  • Tiktoken: A library used to retrieve embeddings from the GPT-2 model.

1. Introduction to Transformers and the Workshop

  • Yufeng Guo, a developer advocate at Google, introduces a workshop on building a transformer model from scratch using JAX and related libraries.
  • The workshop covers the origins, internal structure, and implementation of transformers.
  • The tools used include Flax NNX for model architecture, Optax for loss function and optimizer creation, and Orbax and XLA for accelerated hardware training.

2. Background on Transformers and Embeddings

  • Transformers were initially designed for processing text but have evolved into multi-modal models that can handle images, video, and audio.
  • Embeddings are crucial for converting human-understandable concepts into numerical representations that computers can process.
  • The embedding space is a high-dimensional space where the vocabulary is initially placed randomly.
  • During training, text embeddings are updated to reflect the meaning of words and their relationships.
  • Example: The word "jump" might be close to physical activities and car repair terms in the embedding space.
  • Machine learning models typically process text through the lens of an embedding space rather than raw characters.
  • Resources for learning more about embeddings: Machine Learning Crash Course at developers.google.com and text embeddings from Google Cloud.

3. The "Attention is All You Need" Paper and the Attention Mechanism

  • The transformer architecture is based on the "Attention is All You Need" paper from Google Brain in 2017.
  • The paper introduced the attention mechanism, which calculates and describes the relationship between words in a context.
  • The inputs to attention are three vectors: Query (Q), Key (K), and Value (V).
    • Query (Q): Describes what other words a word cares about.
    • Key (K): Represents what a word wants to present about itself in response to queries.
    • Value (V): Encodes what a word offers if there is a good match.
  • Analogy: Handshake and gift exchange, where the query is someone reaching out, the key is the other hand meeting it, and the value is the gift exchanged.
  • The attention formula includes the softmax function, which normalizes the values coming out of Q.K to sum to 1, highlighting the highest values.
  • Q.K is divided by the square root of D (embedding dimension) to scale down the value passed into softmax, preventing extreme distributions.
  • This process is called scaled dot-product attention.
  • Multi-headed attention involves multiple parallel implementations of attention, allowing the model to attend to different aspects of relationships in the data.

4. Transformer Architecture and Resources

  • The transformer model consists of more than just the attention mechanism.
  • The workshop focuses on the right half of the transformer diagram for building the model.
  • Skip connections allow data from prior to a block to be combined with data after it's been transformed, aiding training.
  • The gray block in the diagram is repeated n times, making the transformer model repetitive.
  • Recommended resources: "The Illustrated Transformer" and chapter four of "How to Scale Your Model" from DeepMind.

5. Setting Up the Development Environment

  • The notebook for the workshop is available in the description and can be run in Kaggle or Colab.
  • In Kaggle, the OpenWebText data set (16 GB) is loaded in the background.
  • In Colab, the OpenWebText data set needs to be manually downloaded (12 GB zip, 18 GB unzipped) and uploaded to Google Drive.
  • Connect the notebook with Weights & Biases (wandb.ai/authorize) to monitor the training run.

6. JAX and Device Configuration

  • jax.devices shows the available devices (TPU v3 on Kaggle, TPU v2 on Colab).
  • The device mesh is set up to distribute the data and model across the available devices.
  • Data parallelism is used, where the data is split across devices, and the model is replicated on each device.
  • The code uses embeddings released with the GPT-2 model, retrieved using the tiktoken library.

7. Hyperparameters and Model Structure (NNX)

  • Hyperparameters are set up to optimize the training run for the available hardware.
  • Batch size: 24 on Colab, 72 on Kaggle, illustrating the difference between TPU v2 and v3.
  • The model structure is defined using the NNX library.
  • The GPT-2 class brings together all the pieces of the model definition.
  • The model consists of an embedding layer, dropout layer, transformer blocks, layer norm, and a final linear layer.
  • The number of transformer blocks is configurable via hyperparameters.
  • The TokenAndPositionEmbedding class combines text embeddings with positional embeddings.
  • Each class inherits from nnx.Module and has an init method and a call method.
  • The TransformerBlock class contains the multi-headed attention mechanism, layer norms, and linear layers.
  • Skip connections are implemented by adding the original inputs to the output of the attention block.

8. Training Loop and Optimization

  • The training loop involves defining a loss function and choosing an optimizer.
  • The loss function computes softmax cross-entropy loss.
  • A single training step computes the gradient of the loss using value_and_grad from NNX.
  • The optimizer (adamw from optax) updates the weights of the model based on the computed gradients.
  • The training loop calls the train_step function and evaluates the model's performance every 200 steps using the validation data set.

9. Accelerators and Infrastructure Options

  • JAX supports CPU, GPU, and TPU.
  • Kaggle offers eight TPU v3 chips for free (up to nine hours).
  • Colab offers eight TPU v2 chips for free, with more recent TPUs and GPUs available with a Pro subscription.
  • GCP allows users to pay for more resources for training large models.
  • Kaggle offers P100 and two T4 GPUs, while Colab offers an NVIDIA T4 GPU for free.
  • JAX is built for distributed parallel processing and allows easy setup of models and data across multiple chips.
  • JAX knows about multiple chips and allows the creation of a device mesh to define parallel strategies.

10. Monitoring Training Progress and Results

  • Training progress can be monitored on Weights & Biases.
  • TPU utilization is shown during training.
  • The training and validation loss decrease steadily.
  • Colab prints outputs showing the progress of training steps.
  • The training run on Colab takes approximately three hours.
  • The training run on Kaggle takes approximately six to seven hours.

11. Model Saving and Loading (Orbax)

  • Orbax is used for saving and loading JAX model checkpoints.
  • The code checks for the existence of relevant file paths and creates a checkpointer.
  • checkpointer.save saves the model state to the desired path.
  • checkpointer.restore restores a model from a checkpoint path.
  • The model architecture is encoded by calling the create_model function.
  • The restored state is put into the model by calling nnx.update.

12. Final Results and Conclusion

  • The model finishes training after a few hours.
  • Training and validation scores are available on Colab and Weights & Biases.
  • Sample outputs are generated.
  • TPU resources are fully utilized.
  • Training the model on GCP with eight TPU v6's (Trillium) takes just 75 minutes with a batch size of 144.
  • The video concludes with thanks and encouragement to continue exploring AI model building and training.

Main Takeaways/Synthesis:

The workshop provides a comprehensive guide to building a transformer model from scratch using JAX and related libraries. It covers the theoretical background of transformers, including embeddings, the attention mechanism, and the overall architecture. The workshop also provides practical guidance on setting up the development environment, configuring JAX for distributed training, defining the model structure using NNX, and implementing the training loop with Optax. Finally, it demonstrates how to save and load the trained model using Orbax. The workshop emphasizes the importance of understanding the underlying concepts and provides resources for further learning. The use of specific tools and libraries, along with detailed explanations, makes this a valuable resource for anyone looking to delve into the world of building and training AI models.

AI summaries can miss context or contain errors. Check important details against the original video.

MAKE IT YOURS

Read. Remember. Reuse.

Free tools

Go a little deeper.

Have a question about this video? Load its transcript to open the video chat.