model = create_model() optimizer = flax.optim.Adam(learning_rate=0.001) @jax.pmap def train_step(optimizer, batch): def loss_fn(model): logits = model(batch['inputs']) # Compute the loss loss = loss_fn(labels=batch['labels'], logits=logits) return loss.mean() # Compute the gradient function grad_fn = jax.grad(loss_fn) # Compute the gradients grad = grad_fn(optimizer.target) # Update the model parameters optimizer = optimizer.apply_gradient(grad) return optimizer for batch in dataset: # Perform a training step optimizer = train_step(optimizer, batch) __ __