Iterable dataset exhausts after a single epoch

Viewed 663

I wanted to train an RNN on the task of sentiment analysis, for this task I was using the IMDB dataset provided by torchtext which contains 50000 movie reviews and it is a python iterator. I used a split=('train', 'test').

I first built a vocab using torchtext.vocab.Vocab and tokenized each of the sentence and then performed numericalisation.

To pad the sequence to the same length I used torch.nn.utils.rnn.pad_sequence and also used a collate_fn together with batch_sampler. Then I loaded the data using torch.utils.data.DataLoader.

The implementation of the RNN Network is fine but the dataloader is exhausted after a single epoch as you can see in the image attached below.

Am I following the right approach to load this iterable-dataset? and why is the dataloader exhausted after a single epoch and how do I overcome this issue.

Pleases refer to the shared colab notebook if you want to see my implementation.

PS. I was following the official changelog of torchtext from github

You can find my implementation here

Dataloader exhausted after a single epoch

1 Answers

The solution is to use torchtext.data.functional.to_map_style_dataset(iter_data) (official doc) to convert your iterable-style dataset to map-style dataset.

Like this:

from torchtext.data.functional import to_map_style_dataset
train_iter = IMDB(split='train')
train_dataset = to_map_style_dataset(train_iter)  #Map-style dataset

and then make a dataloader.

from torch.utils.data import DataLoader
train_dataloader = DataLoader(train_dataset, batch_size=64, collate_fn=collate_fn)

Why is this happening?

I am using above example's naming convention to explain.

The train_iter getting passed to Dataloader is an Iterable-style dataset which means it does not have __getitem__ implemented. It only has __iter__ and __next__ dunders - which makes it a Iterable.

Therefore if I pass an iterable to the Dataloader, the dataloader stops after the StopIteration exception occurs - which will be thrown by __next__ dunder of the iterable-style dataset(train_iter in this case) when the dataset(the iterable) got exhausted.

So we have used the to_map_style_dataset function to convert Iterable-style to map-style dataset. It does so by implementing a __getitem__ dunder and thus Dataloader by default use indices to get items from the dataset.

Another possible way of doing the same thing can also be

If I'll go with iterable-style dataset - I need to create the Dataloader object at every epoch. So after each epoch the new dataloader object will run from start in the for loop.

For better understanding, the differences and use-cases for Iterable-style and Map-style datasets in Pytorch, refer this https://yizhepku.github.io/2020/12/26/dataloader.html

Related