Skip to content

1. The dataset

dataloaders/my_dataset.py → cfg.train_ds_name, valid_ds_name, test_ds_name

def get_ds(ds_params, transform=None)  ->  (dataloader, info)

info["length"] is required. ds_params is that split's *_ds_params. Batch size lives there, not at top level, because splits want different ones.

examples/synthetic/dataloaders/synthetic_dataset.py
import torch
from torch.utils.data import TensorDataset

from gatle_ignite import build_dataloader


def get_ds(ds_params, transform=None):
    """A learnable synthetic task: labels are a fixed linear function of the inputs."""
    n, in_dim = ds_params.get("n", 1024), ds_params.get("in_dim", 64)
    n_classes = ds_params.get("n_classes", 10)

    # The labelling function has its own generator, so train and valid share one task.
    weight_gen = torch.Generator().manual_seed(ds_params.get("task_seed", 1234))
    weight = torch.randn(in_dim, n_classes, generator=weight_gen)

    x_gen = torch.Generator().manual_seed(ds_params.get("seed", 0))
    x = torch.randn(n, in_dim, generator=x_gen)
    y = (x @ weight).argmax(dim=1)

    ds = TensorDataset(x, y)
    return build_dataloader(ds, ds_params), {"length": len(ds)}

build_dataloader(dataset, ds_params, sampler=None) wraps idist.auto_dataloader, except on an exact_sharding split under DDP, which builds a plain DataLoader instead, forcing drop_last=False and ignoring shuffle. It reads:

bs · num_workers · shuffle · drop_last · pin_memory · collate_fn · sampler_params · exact_sharding. Anything else is your dataset's own business. bs is the total across GPUs on every split: under DDP both paths divide it, and num_workers, the way idist.auto_dataloader does.

A sampler or a collate is dotted like every other component:

"sampler_params": {"cls_name": "dataloaders.data_utils.balanced_sampler", "params": {...}}
"collate_fn":     {"cls_name": "dataloaders.data_utils.pad_collate",      "params": {"pad_id": 0}}

naming modules that expose get_sampler(dataset, **params) and get_collate_fn(**params). collate_fn also takes a plain callable, which is what a get_ds usually has in hand.

cfg.aug_name builds one transform and hands it to every split's get_ds. Your dataset decides who actually gets it.

Calling build_dataloader is not compulsory, but skipping it costs you

torch's DistributedSampler either pads an unevenly-divisible eval split or drops its tail. Five samples over two ranks: the default gives [0,2,4] and [1,3,0], scoring sample 0 twice; drop_last=True gives [0,2] and [1,3] and never scores sample 4. Either way every rank agrees on the same wrong number, so the trainer uses ExactDistributedSampler on valid and test. A get_ds that builds its own loader gets none of it, and is warned.

Templates: dataset · sampler · collate · augmentation