class MiniMobileNet(d2l.Classifier):
def __init__(self, arch=((64, 1), (128, 2), (128, 1), (256, 2),
(256, 1), (512, 2), (512, 1)),
lr=0.1, num_classes=10, rngs=None):
super().__init__()
self.save_hyperparameters(ignore=['rngs'])
rngs = nnx.Rngs(d2l.get_key()) if rngs is None else rngs
layers = [nnx.Conv(1, 32, kernel_size=(3, 3), strides=(2, 2),
padding='same', use_bias=False,
kernel_init=nnx.initializers.he_normal(),
rngs=rngs),
nnx.BatchNorm(32, momentum=0.9, rngs=rngs), nnx.relu]
c = 32
for c_out, s in arch:
layers.append(DWSBlock(c, c_out, (s, s), rngs))
c = c_out
layers.extend([lambda x: x.mean(axis=(1, 2)), # global average pooling
nnx.Linear(c, num_classes, rngs=rngs)])
self.net = nnx.Sequential(*layers)
def configure_optimizers(self):
return optax.sgd(self.lr, momentum=0.9)