Getting started with MNIST

By Google for Developers

Share:

Key Concepts

  • JAX Stack: A suite of libraries for machine learning built on JAX.
  • NNX (Neural Network eXtensions): A library within the JAX stack for defining and managing neural network models and their parameters.
  • MNIST: A dataset of handwritten digits, commonly used for machine learning examples.
  • CNN (Convolutional Neural Network): A type of neural network architecture particularly effective for image recognition tasks.
  • Grain: A JAX-native data loading library.
  • Pandas: A Python library for data manipulation and analysis.
  • Parquet: A columnar storage file format.
  • Python Imaging Library (PIL): A library for image manipulation.
  • NumPy: A fundamental library for numerical computation in Python.
  • OptiX: A library for optimizers in JAX.
  • nx.param: A type in NNX used to denote learnable parameters of a model.
  • JIT Compilation (jax.jit): A technique to significantly speed up function execution by compiling them ahead of time.
  • Cross-Entropy Loss: A common loss function for classification tasks.
  • AdamW Optimizer: An optimization algorithm that combines Adam with weight decay.

MNIST Data Loading and Preprocessing

The video demonstrates how to load and preprocess the MNIST dataset for training a CNN using the JAX stack.

  1. Data Download: The MNIST dataset is downloaded from HuggingFace in parquet format. wget commands are provided (commented out) for downloading the train.parquet and test.parquet files.
  2. Pandas Integration: The pandas library is imported to read the parquet files into dataframes named mnest_train_df and mnest_test_df.
  3. Grain Data Loader Setup:
    • A custom Dataset class is defined to wrap the pandas dataframes. This class implements __len__ to return the number of items and __getitem__ to fetch and convert data into NumPy arrays.
    • A helper function convert_to_numpy is used within the Dataset class. This function:
      • Extracts image bytes from a data dictionary.
      • Uses PIL to open the image.
      • Converts the image to a NumPy array.
      • Normalizes pixel values by dividing by 255.
      • Reshapes the array to include a channel dimension (e.g., from (28, 28) to (28, 28, 1)).
      • Extracts the label.
      • Returns both the image and label as NumPy arrays.
    • grain.SequentialSampler is used to iterate through the dataset in order.
    • grain.DataLoader objects are created for both training and testing. These data loaders are configured with:
      • The Dataset object.
      • The Sampler.
      • A grain.batch operation with a specified batch_size and drop_remainder=True to ensure consistent batch sizes.

CNN Model Definition with NNX

A simple Convolutional Neural Network (CNN) model is defined using NNX for MNIST classification.

  1. Architecture: The CNN consists of:
    • Two convolutional layers.
    • Each convolutional layer is followed by average pooling and ReLU activation.
    • The output is flattened.
    • The flattened output is passed through two linear (fully connected) layers.
  2. Instantiation and Visualization:
    • The CNN model is instantiated.
    • nx.display() is used to visualize the model's architecture, which is helpful for debugging.
  3. Dummy Input Test: A dummy input of ones is passed through the model to verify that it runs without errors and produces an output before training.

Optimizer and Metrics Setup

The video details the setup of the optimizer and metrics for training the CNN.

  1. Optimizer:
    • The AdamW optimizer from optix is used.
    • A key update in NNX version 0.11 is highlighted: the wrt=nx.param argument in the optimizer initialization. This explicitly tells the optimizer to only update variables of type nx.param (learnable parameters), ignoring other state variables.
    • nx.display() is used to visualize the optimizer setup.
  2. Metrics: While not explicitly detailed in the transcript, the training and evaluation steps mention updating metrics, implying a metric tracking mechanism is in place (e.g., for loss and accuracy).

Training and Evaluation Steps

The core logic for training and evaluating the model is defined using JAX's JIT compilation.

  1. Loss Function: Cross-entropy loss is calculated between the model's predictions and the true labels.
  2. train_step Function:
    • This function computes the gradients of the loss with respect to the model's parameters using jax.grad.
    • It updates the metrics.
    • It applies the computed gradients to update the model's parameters using the optimizer.
    • It is decorated with nx.jit for JIT compilation.
  3. eval_step Function:
    • This function calculates the loss and updates metrics during evaluation.
    • Crucially, it does not perform gradient updates.
    • It is also decorated with nx.jit for performance.

Training Loop and Execution

The main training loop orchestrates the training and evaluation process.

  1. Iteration: The loop iterates through the training dataset.
  2. Training Step: For each batch, the train_step function is called.
  3. Evaluation and Logging:
    • At specified intervals (controlled by eval_every) or at the end of training, training metrics are computed and recorded.
    • The eval_step function is run on the test dataset.
    • Test metrics are recorded.
    • Metrics are reset after evaluation.
  4. Number of Steps: The process continues for a predefined number of training steps.

Model Prediction and Visualization

After training, the model's performance is assessed on test images.

  1. Evaluation Mode: The model is switched to evaluation mode using model.eval().
  2. Prediction Function (pred_step): A function pred_step is defined to obtain predicted labels from the model.
  3. Batch Prediction: A batch of data is taken from the test dataset, and pred_step is used to make predictions.
  4. Visualization: A grid of test images is plotted, with their corresponding predicted labels displayed. This provides a visual assessment of the model's accuracy.

Further Learning Resources

The video concludes by pointing viewers to additional resources:

  • Coding exercises, quick reference docs, and slides for learning JAX and the JAX AI stack.
  • A YouTube playlist for the "Learning JAX" series.
  • A Discord community for JAX with an invite link.
  • Links to the documentation for JAX, Flax, and the JAX AI stack.
  • A preview of upcoming episodes covering more of the JAX AI stack.

Conclusion/Synthesis

This video provides a practical, step-by-step guide to building and training a CNN for the MNIST dataset using the JAX AI stack, with a particular focus on NNX for model definition and grain for data loading. Key takeaways include the efficient data handling with grain, the clear parameter management in NNX with nx.param, the performance benefits of jax.jit, and the structured approach to defining training and evaluation loops. The example emphasizes best practices for modern JAX development, including explicit parameter handling in optimizers and visualization tools for model architecture.

Chat with this Video

AI-Powered

Load the transcript when you're ready to chat so the initial page stays lighter.

Ready to summarize another video?

Summarize YouTube Video