# Assume train_images, train_labels are loaded NumPy arrays, e.g. from MNIST batch_size = 64 data_loader = build_data_loader(train_images, train_labels, batch_size, seed=42) data_iterator = iter(data_loader) num_steps = 2000 print("Starting training with Grain data pipeline...") running_loss = 0.0 for step in range(num_steps): batch = next(data_iterator) x_batch = jnp.array(batch['image']).reshape(batch_size, -1) y_batch = jnp.array(batch['label']) loss = train_step(model, optimizer, x_batch, y_batch) running_loss += loss.item() if step % 200 == 0 and step > 0: avg_loss = running_loss / 200 print(f"Step {step:4d} | Average Loss: {avg_loss:.4f}") running_loss = 0.0 print("Training complete.") __ __