import optax from flax import nnx # Optimizer Setup tx = optax.adam(learning_rate=0.001) optimizer = nnx.Optimizer(model, tx, wrt=nnx.Param) # Training Step Pattern @nnx.jit def train_step(model, optimizer, batch): def loss_fn(model): logits = model(batch['x']) loss = compute_loss(logits, batch['y']) return loss, logits (loss, logits), grads = nnx.value_and_grad(loss_fn, has_aux=True)(model) optimizer.update(model, grads) return loss # Evaluation Step Pattern @nnx.jit def eval_step(model, batch): logits = model(batch['x']) loss = compute_loss(logits, batch['y']) return loss # Training/Eval Mode model.train() # Enable dropout, stochastic behavior model.eval() # Deterministic inference __ __