How to iterate through composed dataset in pytorch with no overlapped batches?

Viewed 1899

I am looking for a way to connect two DataSets to one, so that it can be trained in one loop. However the batches are not allowed to mix between the datasets. In the following example should only be batches in range 1 to 10 and 41 to 50:

import pandas as pd
import torch
from torch.utils.data import Dataset, DataLoader, ConcatDataset

df1 = pd.DataFrame(list(range(1,11)))
df2 = pd.DataFrame(list(range(41,51)))

class testset(Dataset):
    def __init__(self,data):
        self.data = data

    def __len__(self):
        return len(self.data)
    
    def __getitem__(self, index):
        return self.data[0][index]

testdataset1 = testset(df1)
testdataset2 = testset(df2)

datasets = []
datasets.append(testdataset1)
datasets.append(testdataset2)

concat_dataset = ConcatDataset(datasets)

loader = DataLoader(
    concat_dataset,
    shuffle=False,
    num_workers=0,
    batch_size=3
)

for data in loader:
    print(data)

tensor([1, 2, 3])

tensor([4, 5, 6])

tensor([7, 8, 9])

tensor([10, 41, 42]) ← That should not exist

tensor([43, 44, 45])

tensor([46, 47, 48])

tensor([49, 50])

In the real case I am combining two timeseries, where overlapping in batches with values of both datasets causes a littlebit trouble…

This shouldn’t be a though one, right?

1 Answers

If you do that you are not creating random batches anymore (these are pseudo-random) as batch elements are restricted (if the first element comes from 0 dataset, rest of them also have to).

Short description:

  • batch_size has to be specified (as sample generation is dependent on it)
  • Optional length argument as now this dataset can be of any length (sample is taken from some dataset via modulo operation)
  • Start with 0th dataset and generate batch from it
  • Move to another dataset (you can switch within __getitem__ method):
    • randomly: method _new_random_dataset
    • simply next one: method _next_dataset

Below is a torch.utils.data.Dataset custom instance which does what you want:

class Merger(torch.utils.data.Dataset):
    def __init__(
        self, *datasets: torch.utils.data.Dataset, batch_size: int, length: int = None
    ):
        self.datasets = datasets
        self.batch_size = batch_size

        if length is None:
            self._len = sum(len(d) for d in self.datasets)
        else:
            self._len = length

        # Keep in internal var how many items we've generated
        # Only possible dataset switch when new batch is created
        self._items_generated = 0
        # First batch will always go from the 0th dataset
        self._dataset_index = 0

    def __len__(self):
        return self._len

    def _next_dataset(self):
        if self._dataset_index == len(self.datasets) - 1:
            self._dataset_index = 0
        else:
            self._dataset_index += 1

    def _new_random_dataset(self):
        self._dataset_index = random.randrange(0, len(self.datasets))

    def __getitem__(self, index):
        if self._items_generated >= self.batch_size:
            self._items_generated = 0
            # self._next_dataset()
            self._new_random_dataset()

        self._items_generated += 1
        return self.datasets[self._dataset_index][
            index % len(self.datasets[self._dataset_index])
        ]

And example usage for you to verify:

df1 = pd.DataFrame(list(range(1, 11)))
df2 = pd.DataFrame(list(range(41, 51)))

ds = Merger(testset(df1), testset(df2), batch_size=3)

loader = torch.utils.data.DataLoader(ds, shuffle=False, num_workers=0, batch_size=3)

for data in loader:
    print(data)
Related