add mnist_classifier/ (checkpoint + code)
Browse files
mnist_classifier/classifier.py
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import jax
|
| 2 |
+
from flax import linen as nn
|
| 3 |
+
|
| 4 |
+
Array = jax.Array
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
class Classifier(nn.Module):
|
| 8 |
+
num_classes: int
|
| 9 |
+
base_channels: int = 16
|
| 10 |
+
|
| 11 |
+
@nn.compact
|
| 12 |
+
def __call__(self, x: Array) -> Array:
|
| 13 |
+
h = nn.Conv(self.base_channels, kernel_size=(3, 3))(x)
|
| 14 |
+
h = nn.relu(h)
|
| 15 |
+
h = nn.max_pool(h, window_shape=(2, 2), strides=(2, 2))
|
| 16 |
+
h = nn.Conv(self.base_channels * 2, kernel_size=(3, 3))(h)
|
| 17 |
+
h = nn.relu(h)
|
| 18 |
+
h = nn.max_pool(h, window_shape=(2, 2), strides=(2, 2))
|
| 19 |
+
h = h.reshape(h.shape[0], -1)
|
| 20 |
+
return nn.Dense(self.num_classes)(h)
|
mnist_classifier/load.py
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import pickle
|
| 2 |
+
|
| 3 |
+
from classifier import Classifier
|
| 4 |
+
|
| 5 |
+
with open("mnist_classifier.pkl", "rb") as f:
|
| 6 |
+
bundle = pickle.load(f)
|
| 7 |
+
meta = bundle["meta"]
|
| 8 |
+
params = bundle["params"]
|
| 9 |
+
|
| 10 |
+
model = Classifier(
|
| 11 |
+
num_classes=meta["num_classes"],
|
| 12 |
+
base_channels=meta["base_channels"],
|
| 13 |
+
)
|
mnist_classifier/mnist_classifier.pkl
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:2f68078916e0751319df3dd4280be0293bdbeff895a948b724cd9ded776846aa
|
| 3 |
+
size 101781
|