Is it advisable to use the same torch Dataset class for training and predicting?

Viewed 218

I have recently started using PyTorch and I liked it for its object-oriented style. However, I wonder what’s the best and advised workflow when predicting the model. I wanted to use a custom Dataset class I wrote and which I use for training and validating my model. This class is a map-style dataset, therefore I implement __getitem__ method to return image and target:

class CustomDataset:

    def __init__(self, ...):
        ...

    def __getitem__(self, image_id):
        ....
        return (
            torch.tensor(image, dtype=torch.float),
            torch.tensor(target, dtype=torch.long),
       )

However, when I’m using this class for predicting I don’t have any targets to return. My current workaround is something like

def __getitem__(self, image_id):
    ....
    if predict:
        return (
            torch.tensor(image, dtype=torch.float),
            np.nan,
       )
   else:
        return (
             torch.tensor(image, dtype=torch.float),
             torch.tensor(target, dtype=torch.long),
       )

However, I wonder if there’s a better way to do it. And at the same time, as it feels a bit unnatural, I started wondering if it is even advisable to use the same class for training and predicting (it should be, but the clunkiness of my solutions makes me wonder). Of course, I could not return a tuple at all, but only a first element, but this still needs if-else.

3 Answers

PyTorch's DataSet class is really simple. So, do not overthink it. It's not much more than a wrapper for accessing your data.

You don't have to return a tuple, not even Tensors. You can return whatever data you want. Commonly, it will be in one of those styles:

  • For unsupervised data: Sample or (Sample, None)
  • For supervised data: (Sample, Label)
  • For supervised data with multiple targets, e.g. object detection: (Sample, [Label1, Label2, ...]) or (Sample, Label1, Label2, ...)

It is also common to use the same DataSet class for train / test.

So, in your case, simply return the sample or a tuple (sample, None) as done in torchvision and adjust your pipeline accordingly. I'd not suggest using np.nan as it would fail a simple None check (np.nan == None). Also, I'd encourage you to inherit from torch.data.Dataset.

If however, your pipeline forces you to use a tuple or has other constraints I'd suggest to rephrase your question.

You must write code to create a Dataset that matches your data and problem scenario; no two Dataset implementations are exactly the same. On the other hand, a DataLoader object is used mostly the same no matter which Dataset object it's associated with. For example:

class MyDataSet(T.utils.data.Dataset):
  # implement custom code to load data here

my_ds = MyDataset("my_train_data.txt")
my_ldr = torch.utils.data.DataLoader(my_ds, 10, True)
for (idx, batch) in enumerate(my_ldr):
  . . .

I think if you want a "pure" (=no if) solution, you could define a "don't care" class, and have your losses ignore it in some way (probably done with masking internally, it is technically an "if", but it is vectorized).

For example, see CrossEntropyLoss has ignore_index to handle such cases. This makes me believe don't care class index is the by-design way to go.

class CDiscountDataset:
    self.ignore_index = self._preprocess_number_of_classes() + 1

    def __getitem__(self, image_id):
        target_tensor = torch.tensor(target, dtype=torch.long)

        if predict:
            return (
                torch.tensor(image, dtype=torch.float),
                torch.ones_like(target_tensor, dtype=torch.long) * self.ignore_index,
            )
        else:
            return (
                torch.tensor(image, dtype=torch.float),
                target_tensor,
           )

As a side note, I am using Pytorch-lightning which gives nice abstractions to the training pipeline, and it implicitly assumes only tuples of tensors return from the dataloader, which makes me believe returning other types is "less canonical", which strengthens my belief that the above "don't care" way is the way to go.

Related