# Training configuration num_epochs = 5 train_steps_per_epoch = 1000 eval_steps = 200 print("Starting training...") print("=" * 60) for epoch in range(num_epochs): # Training model.train() # Set model to training mode train_loss, train_acc = 0.0, 0.0 for step, batch in enumerate(train_ds.as_numpy_iterator()): if step >= train_steps_per_epoch: break # Convert to JAX arrays batch = {k: jnp.array(v) for k, v in batch.items()} loss, acc = train_step(model, optimizer, batch) train_loss += loss train_acc += acc train_loss /= train_steps_per_epoch train_acc /= train_steps_per_epoch # Evaluation model.eval() # Set model to evaluation mode eval_loss, eval_acc = 0.0, 0.0 eval_batches = 0 for step, batch in enumerate(test_ds.as_numpy_iterator()): if step >= eval_steps: break batch = {k: jnp.array(v) for k, v in batch.items()} loss, acc = eval_step(model, batch) eval_loss += loss eval_acc += acc eval_batches += 1 eval_loss /= eval_batches eval_acc /= eval_batches print(f"Epoch {epoch + 1}/{num_epochs}") print(f"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f}") print(f"Eval Loss: {eval_loss:.4f} | Eval Acc: {eval_acc:.4f}") print("Training complete!") __ __