A lightweight MNIST digit classification model implemented from scratch in JAX.
Open
http://localhost:8501 in your browser, draw a digit, and press
Predict.
1from huggingface_hub import hf_hub_download
2import pickle
3import jax
4from mlp_model import MLP, activation_map
5
6# Download parameters
7path = hf_hub_download(repo_id="mysticdan/mlp-mnist", filename="mlp_mnist_model.pkl")
8with open(path, "rb") as f:
9 data = pickle.load(f)
10
11# Reconstruct MLP
12mlp = MLP.__new__(MLP)
13mlp.params = data["params"]
14mlp.activation = tuple(
15 activation_map["identity"] if name == "<lambda>" else activation_map[name]
16 for name in data["activation_names"]
17)
18mlp.key = jax.random.key(0)
19
20# Inference
21import jax.numpy as jnp
22# assume `sample` is a (784,) float32 array normalized [0,1]
23logits = mlp.predict(sample[None, :])
24pred = int(jnp.argmax(logits, axis=1)[0])
25print("Predicted digit:", pred)