class MLP(nnx.Module):
"""A three-hidden-layer ReLU network under standard parametrization."""
def __init__(self, width, rngs=None):
rngs = nnx.Rngs(0) if rngs is None else rngs
init = nnx.initializers.variance_scaling(1.0, 'fan_in', 'normal')
self.fc_in = nnx.Linear(784, width, kernel_init=init, rngs=rngs)
self.fc_h1 = nnx.Linear(width, width, kernel_init=init, rngs=rngs)
self.fc_h2 = nnx.Linear(width, width, kernel_init=init, rngs=rngs)
self.fc_out = nnx.Linear(width, 10, kernel_init=init, rngs=rngs)
def features(self, X):
h = nnx.relu(self.fc_h1(nnx.relu(self.fc_in(X))))
return nnx.relu(self.fc_h2(h))
def __call__(self, X):
return self.fc_out(self.features(X))
def configure_adam(self, lr):
return nnx.Optimizer(self, optax.adam(lr), wrt=nnx.Param)