class ArchiveDigitizer(nnx.Module): def __init__(self, hidden_dim: int, num_classes: int, rngs: nnx.Rngs): self.linear1 = nnx.Linear(784, hidden_dim, rngs=rngs) self.linear2 = nnx.Linear(hidden_dim, num_classes, rngs=rngs) def __call__(self, x): x = self.linear1(x) x = nnx.relu(x) return self.linear2(x) rngs = nnx.Rngs(jax.random.PRNGKey(42)) model = ArchiveDigitizer(hidden_dim=128, num_classes=10, rngs=rngs) tx = optax.adam(learning_rate=0.005) optimizer = nnx.Optimizer(model, tx=tx, wrt=nnx.Param) @nnx.jit def train_step(model, optimizer, batch_images, batch_labels): def loss_fn(m): logits = m(batch_images) loss = optax.softmax_cross_entropy_with_integer_labels( logits=logits, labels=batch_labels ).mean() return loss loss_val, grads = nnx.value_and_grad(loss_fn)(model) optimizer.update(model, grads) return loss_val __ __