import jax from jax import numpy as jnp # Define a model using FLAX with functional programming style @jax.jit def model(params, inputs): dense = flax.nn.Dense(inputs.shape[-1], features=64) x = dense.initialize_carry(jax.random.PRNGKey(0), inputs) x = flax.nn.relu(x) output = flax.nn.Dense(x.shape[-1], features=10).initialize_carry(jax.random.PRNGKey(0), x) return output # Train the model optimizer = flax.optim.Adam(learning_rate=0.001).create(model.params) for batch in dataset: optimizer = optimizer.train_step(batch) __ __