import optax from flax.training import train_state class TrainState(train_state.TrainState): pass # No additional attributes needed for now # Define the optimizer learning_rate = 0.001 optimizer = optax.adam(learning_rate) # Initialize the training state state = TrainState.create(apply_fn=model.apply, params=params, tx=optimizer) __ __