import jax import jax.numpy as jnp from flax import nnx import optax import tensorflow_datasets as tfds import tensorflow as tf # Disable TensorFlow GPU to avoid conflicts with JAX tf.config.set_visible_devices([], 'GPU') # Load MNIST def load_mnist(batch_size=32): """Load and preprocess MNIST dataset.""" def preprocess(sample): image = tf.cast(sample['image'], tf.float32) / 255.0 label = sample['label'] return {'image': image, 'label': label} train_ds = tfds.load('mnist', split='train') test_ds = tfds.load('mnist', split='test') train_ds = (train_ds .map(preprocess) .shuffle(1024) .batch(batch_size, drop_remainder=True) .prefetch(1)) test_ds = (test_ds .map(preprocess) .batch(batch_size, drop_remainder=True) .prefetch(1)) return train_ds, test_ds train_ds, test_ds = load_mnist(batch_size=32) __ __