Coding a Neural Network from Scratch in Pure JAX - Machine Learning with JAX - Tutorial 3
Aleksa Gordić - The AI Epiphany via YouTube
You’re only 3 weeks away from a new language
AI, Data Science & Cloud Certificates from Google, IBM & Meta
Overview
Google, IBM & Meta Certificates – 40% Off
One Coursera Plus subscription covers most Professional Certificates on Coursera.
Unlock All Certificates
This tutorial builds and trains a multi-layer perceptron classifier on MNIST using pure JAX, with PyTorch data loaders. It also covers weight and embedding visualization with t-SNE and analysis of dead neurons.
Syllabus
Intro, structuring the code
MLP initialization function
Prediction function
PyTorch MNIST dataset
PyTorch data loaders
Training loop
Adding the accuracy metric
Visualize the image and prediction
Small code refactoring
Visualizing MLP weights
Visualizing embeddings using t-SNE
Analyzing dead neurons
Outro
Taught by
Aleksa Gordić - The AI Epiphany