# Create and train a PyTorch model for digit classification using the MNIST dataset

## In this learning path

- [Introduction](https://learn.arm.com/learning-paths/cross-platform/pytorch-digit-classification-arch-training/)
- [Prepare a PyTorch Development Environment](https://learn.arm.com/learning-paths/cross-platform/pytorch-digit-classification-arch-training/intro/)
- [Create a PyTorch model for MNIST](https://learn.arm.com/learning-paths/cross-platform/pytorch-digit-classification-arch-training/model/)
- [About PyTorch Model Training](https://learn.arm.com/learning-paths/cross-platform/pytorch-digit-classification-arch-training/intro2/)
- [Perform Training and Save the Model](https://learn.arm.com/learning-paths/cross-platform/pytorch-digit-classification-arch-training/datasets-and-training/)
- [Deploy the Model for Inference](https://learn.arm.com/learning-paths/cross-platform/pytorch-digit-classification-arch-training/inference/)
- [Learn about Inference on Android](https://learn.arm.com/learning-paths/cross-platform/pytorch-digit-classification-arch-training/intro-android/)
- [Create an Android Application](https://learn.arm.com/learning-paths/cross-platform/pytorch-digit-classification-arch-training/user-interface/)
- [Prepare the Test Data](https://learn.arm.com/learning-paths/cross-platform/pytorch-digit-classification-arch-training/prepare-data/)
- [Run the Application](https://learn.arm.com/learning-paths/cross-platform/pytorch-digit-classification-arch-training/app/)
- [Optimizing Neural Network Models in PyTorch](https://learn.arm.com/learning-paths/cross-platform/pytorch-digit-classification-arch-training/intro-opt/)
- [Create an optimized PyTorch model for MNIST](https://learn.arm.com/learning-paths/cross-platform/pytorch-digit-classification-arch-training/model-opt/)
- [Run optimization](https://learn.arm.com/learning-paths/cross-platform/pytorch-digit-classification-arch-training/optimisation/)
- [Update the Android application](https://learn.arm.com/learning-paths/cross-platform/pytorch-digit-classification-arch-training/mobile-app/)
- [Next Steps](https://learn.arm.com/learning-paths/cross-platform/pytorch-digit-classification-arch-training/_next-steps/)

## About this Learning Path

| Skill level:  | Advanced            |
|---------------|---------------------|
| Reading time: | 2 hrs 40 min        |
| Last updated: | 03 Aug 2026         |

| Author:            | Dawid Borycki [GitHub](https://github.com/dawidborycki)    |
|---------------------|------------------------------------------------------------|
| Arm IP:             | [Cortex-A](https://support.arm.com/?tab=compute-ip&Product%20Type=Application%20Processors) [Neoverse](https://support.arm.com/?tab=compute-ip&Product%20Type=Infrastructure%20Processors) |
| Tags:               | [ML](https://learn.arm.com/tag/ml) [Windows](https://learn.arm.com/tag/windows) [Linux](https://learn.arm.com/tag/linux) [macOS](https://learn.arm.com/tag/macos) [Android Studio](https://learn.arm.com/tag/android-studio) [Visual Studio Code](https://learn.arm.com/tag/visual-studio-code) |

### Who is this for?
This is an advanced topic for software developers interested in learning how to use PyTorch to create and train a feedforward neural network for digit classification, and also software developers interested in learning how to use and apply optimizations to the trained model in an Android application.

### What will you learn?
Upon completion of this Learning Path, you will be able to:
- Prepare a PyTorch development environment.
- Download and prepare the MNIST dataset.
- Create and train a neural network architecture using PyTorch.
- Create an Android app and load the pre-trained model.
- Prepare an input dataset.
- Measure the inference time.
- Optimize a neural network architecture using quantization and fusing.
- Deploy an optimized model in an Android application.

### Prerequisites
Before starting, you will need the following:
- A machine that can run Python3, Visual Studio Code, and Android Studio.
- For the OS, you can use Windows, Linux, or macOS.

### Summary
You’ll build and train a PyTorch feedforward neural network for MNIST digit classification, then use it for inference and Android deployment. You’ll set up Python, prepare the dataset, and train the model. Then, you’ll reload the model’s saved parameters and apply the same preprocessing to new images. You’ll quantize and fuse an optimized variant, integrate it into an Android application, and compare inference times before deployment.

### Frequently asked questions
<details>
<summary>How do I know the MNIST dataset is set up correctly before training?</summary>
Confirm that the training and test splits download without errors and that the `DataLoader` returns batches. With `batch_size` set to 32, the first batch contains 32 images of 28x28 pixels and corresponding labels from 0–9.
</details>

<details>
<summary>What should I look for during training to confirm the model is learning?</summary>
Monitor the loss and confirm that it decreases over epochs. Evaluate on the test data periodically to check whether predictions align with the true labels more often.
</details>

<details>
<summary>After training, what do I need to load the model for inference?</summary>
Use the model file produced during training and provide its path when loading. Recreate the same model architecture before loading the saved parameters.
</details>

<details>
<summary>How do I validate that inference works on new images?</summary>
Apply the training preprocessing, including normalization and tensor conversion, then run a prediction. The output maps to a digit from 0–9. Compare it with a known test-set label.
</details>

<details>
<summary>How do I compare unoptimized and optimized models after quantization and fusing?</summary>
Measure inference time with the original model, then repeat the measurement after quantization and fusing with identical inputs and conditions. Use the timings to choose the model for the Android application.
</details>
