import grain.python as grain import numpy as np import jax import jax.numpy as jnp from flax import nnx import optax class ArrayDataSource(grain.RandomAccessDataSource): """Wraps in-memory NumPy arrays as a Grain data source.""" def __init__(self, images, labels): self._images = images self._labels = labels def __len__(self): return len(self._images) def __getitem__(self, index): return { 'image': self._images[index], 'label': self._labels[index], } class Normalize(grain.MapTransform): def map(self, element): element['image'] = element['image'].astype(np.float32) / 255.0 return element def build_data_loader(images, labels, batch_size, seed): source = ArrayDataSource(images, labels) sampler = grain.IndexSampler( num_records=len(source), shuffle=True, num_epochs=None, seed=seed, ) operations = [ Normalize(), grain.Batch(batch_size=batch_size, drop_remainder=True), ] return grain.DataLoader( data_source=source, sampler=sampler, operations=operations, worker_count=4, ) __ __