# Get the full optimizer state opt_state = nnx.state(optimizer) # The state is a nested structure containing: # - Step count # - Momentum buffers (for Adam's first moment) # - Velocity buffers (for Adam's second moment) # - Any other state from transformations in the chain print(jax.tree_util.tree_map(lambda x: x.shape, opt_state)) __ __