Project-MONAI / Project-MONAI/MONAI

Allow `DataLoader` and `Dataset` to retain `Generic` features from torch

Open
#8,333 3 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
8.7k
Forks
1.6k
Avg merge
5d 1h
Merged PRs (30d)
20

Description

Is your feature request related to a problem? Please describe.
I would like to be able to provide better type hints for my monai code. The DataLoader and Dataset classes in monai inherit from torch but hide the fact that in torch these are generic classes in torch. For example, in torch I can define a dataset like:

from torch.utils.data import Dataset

class MyData(TypedDict):
    filename: str
    image: torch.Tensor
    segmentation: torch.Tensor

my_dataset: Dataset[MyData] = create_data()

This means that elsewhere in the code I can have a better idea what the data will look like. This kind of thing isn't possible if I'm using the monai code.

Describe the solution you'd like
I think you could solve this by doing something like the following (for Dataset - you'd need to do something similar in DataLoader):

import collections.abc
from typing import Any, Mapping, Sequence, TypeVar, Union, overload
import numpy as np
import torch
from torch.utils.data import Dataset as _TorchDataset
from torch.utils.data import Subset as _TorchSubset


NdarrayOrTensor = Union[np.ndarray, torch.Tensor]


T = TypeVar(
    "T",
    bound=NdarrayOrTensor | Sequence[NdarrayOrTensor] | Mapping[Any, NdarrayOrTensor],
)
class Dataset(_TorchDataset[T]):

    # Leave the rest of the class as-is
    ...

    @overload
    def __getitem__(self, index: slice) -> _TorchSubset[T]:
        ...
    @overload
    def __getitem__(self, index: Sequence[int]) -> _TorchSubset[T]:
        ...
    @overload
    def __getitem__(self, index: int) -> T:
        ...
    
    def __getitem__(self, index: int | slice | Sequence[int]) -> T | _TorchSubset[T]:
        """
        Returns a `Subset` if `index` is a slice or Sequence, a data item otherwise.
        """
        if isinstance(index, slice):
            # dataset[:42]
            start, stop, step = index.indices(len(self))
            indices = range(start, stop, step)
            return _TorchSubset(dataset=self, indices=indices)
        if isinstance(index, collections.abc.Sequence):
            # dataset[[1, 3, 4]]
            return _TorchSubset(dataset=self, indices=index)
        return self._transform(index)

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Read the MONAI definitions of DataLoader and Dataset and compare their inheritance and getitem behavior with torch's generic classes. Confirm the typing approach supports annotations such as Dataset[MyData] and retains existing Dataset and DataLoader behavior, including subset indexing.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
data, machine-learning
Issue type
Feature
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.