torchrecurrent.benchmarks.sequential_mnist#
- torchrecurrent.benchmarks.sequential_mnist(images, targets, permutation=None, return_dataloader=True, batch_size=64, shuffle=True, *, normalize=True, dtype=None, device=None, generator=None, **dataloader_kwargs)[source]#
Convert MNIST tensors into the sequential or permuted-MNIST task.
Every 28 by 28 image becomes a sequence of 784 scalar inputs in row-major raster order. With a permutation, the same fixed pixel ordering is applied to every image, matching the pixel-by-pixel and permuted task setup of Arjovsky et al. (2016), Section 5.3 (https://proceedings.mlr.press/v48/arjovsky16.html); the specific base scan direction is an arbitrary fixed convention and does not affect task difficulty. The caller supplies the MNIST tensors so this package does not require a dataset-download dependency.
- Parameters:
images – MNIST images with shape
(N, 28, 28)or(N, 1, 28, 28).targets – Digit class indices with shape
(N,).permutation – Optional permutation of the integers from 0 through 783. Reuse the same tensor for training and test data.
return_dataloader – Return a data loader when true, otherwise tensors.
batch_size – Batch size of the returned data loader.
shuffle – Whether the returned data loader shuffles samples.
normalize – Convert integer pixels from
[0, 255]to[0, 1]. Floating-point inputs are assumed to be normalized already.dtype – Floating-point dtype of the returned inputs. Defaults to the current PyTorch default dtype.
device – Device on which to place inputs and targets.
generator – Optional generator used by the data loader when shuffling.
**dataloader_kwargs – Additional arguments passed to
torch.utils.data.DataLoader.
- Returns:
A data loader, or
(sequences, targets)whenreturn_dataloader=False. Sequences have shape(N, 784, 1)and targets have shape(N,)with dtypetorch.long.