Limitation
AugmentedDataLoader (braindecode/augmentation/base.py) applies transforms in place: its collate runs default_collate(batch) then transform(X, y) and returns a same-size batch. So augmentation is stochastic-per-epoch only — there's no way to express a fixed expansion (keep the originals and append N augmented copies), which several EEG augmentations are defined as.
Example: the EEG-Inception MI augmentation (Zhang 2021) builds a 6× training set (1 original + 5 augmented). With probability=1.0 the current loader augments every sample and the model never sees clean data — both unfaithful and, empirically, harmful (it degrades the strongest subjects).
Suggested feature
Add an n_augmentation (or multiply) kwarg: when > 0, the collate returns 1 original + n_augmentation transformed copies (6× for n_augmentation=5), keeping the clean originals. Fully backwards-compatible (default 0 = current behaviour). Sketch:
def __init__(self, dataset, transforms=None, device=None, n_augmentation=0, **kwargs):
super().__init__(...)
if n_augmentation > 0:
base = self.collate_fn
def grow(batch):
aug = [base(batch) for _ in range(n_augmentation)]
clean = default_collate(batch)
dev = aug[0][0].device
xs = [clean[0].to(dev)] + [a[0] for a in aug]
ys = [clean[1].to(dev)] + [a[1] for a in aug]
return torch.cat(xs), torch.cat(ys)
self.collate_fn = grow
Happy to PR.
Limitation
AugmentedDataLoader(braindecode/augmentation/base.py) applies transforms in place: its collate runsdefault_collate(batch)thentransform(X, y)and returns a same-size batch. So augmentation is stochastic-per-epoch only — there's no way to express a fixed expansion (keep the originals and append N augmented copies), which several EEG augmentations are defined as.Example: the EEG-Inception MI augmentation (Zhang 2021) builds a 6× training set (1 original + 5 augmented). With
probability=1.0the current loader augments every sample and the model never sees clean data — both unfaithful and, empirically, harmful (it degrades the strongest subjects).Suggested feature
Add an
n_augmentation(ormultiply) kwarg: when > 0, the collate returns1 original + n_augmentationtransformed copies (6× forn_augmentation=5), keeping the clean originals. Fully backwards-compatible (default 0 = current behaviour). Sketch:Happy to PR.