Getting started with MNIST
By Google for Developers
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.
- Data Download: The MNIST dataset is downloaded from HuggingFace in
parquetformat.wgetcommands are provided (commented out) for downloading thetrain.parquetandtest.parquetfiles. - Pandas Integration: The
pandaslibrary is imported to read theparquetfiles into dataframes namedmnest_train_dfandmnest_test_df. - Grain Data Loader Setup:
- A custom
Datasetclass 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_numpyis used within theDatasetclass. 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.SequentialSampleris used to iterate through the dataset in order.grain.DataLoaderobjects are created for both training and testing. These data loaders are configured with:- The
Datasetobject. - The
Sampler. - A
grain.batchoperation with a specifiedbatch_sizeanddrop_remainder=Trueto ensure consistent batch sizes.
- The
- A custom
CNN Model Definition with NNX
A simple Convolutional Neural Network (CNN) model is defined using NNX for MNIST classification.
- 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.
- Instantiation and Visualization:
- The CNN model is instantiated.
nx.display()is used to visualize the model's architecture, which is helpful for debugging.
- 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.
- Optimizer:
- The
AdamWoptimizer fromoptixis used. - A key update in NNX version 0.11 is highlighted: the
wrt=nx.paramargument in the optimizer initialization. This explicitly tells the optimizer to only update variables of typenx.param(learnable parameters), ignoring other state variables. nx.display()is used to visualize the optimizer setup.
- The
- 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.
- Loss Function: Cross-entropy loss is calculated between the model's predictions and the true labels.
train_stepFunction:- 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.jitfor JIT compilation.
- This function computes the gradients of the loss with respect to the model's parameters using
eval_stepFunction:- This function calculates the loss and updates metrics during evaluation.
- Crucially, it does not perform gradient updates.
- It is also decorated with
nx.jitfor performance.
Training Loop and Execution
The main training loop orchestrates the training and evaluation process.
- Iteration: The loop iterates through the training dataset.
- Training Step: For each batch, the
train_stepfunction is called. - Evaluation and Logging:
- At specified intervals (controlled by
eval_every) or at the end of training, training metrics are computed and recorded. - The
eval_stepfunction is run on the test dataset. - Test metrics are recorded.
- Metrics are reset after evaluation.
- At specified intervals (controlled by
- 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.
- Evaluation Mode: The model is switched to evaluation mode using
model.eval(). - Prediction Function (
pred_step): A functionpred_stepis defined to obtain predicted labels from the model. - Batch Prediction: A batch of data is taken from the test dataset, and
pred_stepis used to make predictions. - 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-PoweredLoad the transcript when you're ready to chat so the initial page stays lighter.