import grain.python as grain import numpy as np # DataSource: raw access to records class MNISTSource(grain.RandomAccessDataSource): 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], } source = MNISTSource(images=train_images, labels=train_labels) # Sampler: order and reproducibility sampler = grain.IndexSampler( num_records=len(source), shuffle=True, num_epochs=None, seed=42, ) # Transformations: processing pipeline class Normalize(grain.MapTransform): def map(self, element): element['image'] = element['image'].astype(np.float32) / 255.0 return element operations = [ Normalize(), grain.Batch(batch_size=64, drop_remainder=True), ] # Assemble the DataLoader data_loader = grain.DataLoader( data_source=source, sampler=sampler, operations=operations, worker_count=4, ) # Use it as an iterator data_iterator = iter(data_loader) first_batch = next(data_iterator) print(f"Batch image shape: {first_batch['image'].shape}") print(f"Batch label shape: {first_batch['label'].shape}") __ __