# Debugging: eager execution, full pdb access with jax.disable_jit(): debug_output = train_step(model, optimizer, tiny_batch) # Production: full speed output = train_step(model, optimizer, real_batch) __ __