# Define an Adam optimizer with learning rate 0.001 optimizer = flax.optim.Adam(learning_rate=0.001) # Define the loss function loss_fn = flax.nn.logits_cross_entropy_loss 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) __ __