@d2l.add_to_class(FashionMNIST)
def get_dataloader(self, train):
data = self.train if train else self.val
process = lambda X, y: (tf.expand_dims(X, axis=3) / 255,
tf.cast(y, dtype='int32'))
resize_fn = lambda X, y: (tf.image.resize_with_pad(X, *self.resize), y)
shuffle_buf = len(data[0]) if train else 1
# `drop_remainder=train` keeps every training minibatch the same
# shape, so JAX does not retrace the `@jax.jit`'d step function for
# a smaller last batch.
dataset = (tf.data.Dataset.from_tensor_slices(process(*data)).shuffle(
shuffle_buf).batch(self.batch_size, drop_remainder=train).map(
resize_fn))
return d2l.TensorFlowDataLoader(dataset)