torch_em.data.concat_dataset

 1import numpy as np
 2
 3from torch.utils.data import Dataset, Subset
 4
 5
 6class ConcatDataset(Dataset):
 7    """Dataset to concatenate multiple PyTorch datasets.
 8
 9    Args:
10        datasets: The datasets to concatenate.
11    """
12    def __init__(self, *datasets: Dataset):
13        self.datasets = datasets
14        dataset = datasets[0]
15        while isinstance(dataset, Subset):
16            dataset = dataset.dataset
17        self.ndim = dataset.ndim
18
19        # compute the number of samples for each volume
20        self.ds_lens = [len(dataset) for dataset in self.datasets]
21        self._len = sum(self.ds_lens)
22
23        # compute the offsets for the samples
24        self.ds_offsets = np.cumsum(self.ds_lens)
25
26    def __len__(self):
27        return self._len
28
29    def __getitem__(self, idx):
30        # find the dataset id corresponding to this index
31        ds_idx = 0
32        while True:
33            if idx < self.ds_offsets[ds_idx]:
34                break
35            ds_idx += 1
36
37        # get sample from the dataset
38        ds = self.datasets[ds_idx]
39        offset = self.ds_offsets[ds_idx - 1] if ds_idx > 0 else 0
40        idx_in_ds = idx - offset
41        assert idx_in_ds < len(ds) and idx_in_ds >= 0, f"Failed with: {idx_in_ds}, {len(ds)}"
42        return ds[idx_in_ds]
class ConcatDataset(typing.Generic[+_T_co]):
 7class ConcatDataset(Dataset):
 8    """Dataset to concatenate multiple PyTorch datasets.
 9
10    Args:
11        datasets: The datasets to concatenate.
12    """
13    def __init__(self, *datasets: Dataset):
14        self.datasets = datasets
15        dataset = datasets[0]
16        while isinstance(dataset, Subset):
17            dataset = dataset.dataset
18        self.ndim = dataset.ndim
19
20        # compute the number of samples for each volume
21        self.ds_lens = [len(dataset) for dataset in self.datasets]
22        self._len = sum(self.ds_lens)
23
24        # compute the offsets for the samples
25        self.ds_offsets = np.cumsum(self.ds_lens)
26
27    def __len__(self):
28        return self._len
29
30    def __getitem__(self, idx):
31        # find the dataset id corresponding to this index
32        ds_idx = 0
33        while True:
34            if idx < self.ds_offsets[ds_idx]:
35                break
36            ds_idx += 1
37
38        # get sample from the dataset
39        ds = self.datasets[ds_idx]
40        offset = self.ds_offsets[ds_idx - 1] if ds_idx > 0 else 0
41        idx_in_ds = idx - offset
42        assert idx_in_ds < len(ds) and idx_in_ds >= 0, f"Failed with: {idx_in_ds}, {len(ds)}"
43        return ds[idx_in_ds]

Dataset to concatenate multiple PyTorch datasets.

Arguments:
  • datasets: The datasets to concatenate.
ConcatDataset(*datasets: torch.utils.data.dataset.Dataset)
13    def __init__(self, *datasets: Dataset):
14        self.datasets = datasets
15        dataset = datasets[0]
16        while isinstance(dataset, Subset):
17            dataset = dataset.dataset
18        self.ndim = dataset.ndim
19
20        # compute the number of samples for each volume
21        self.ds_lens = [len(dataset) for dataset in self.datasets]
22        self._len = sum(self.ds_lens)
23
24        # compute the offsets for the samples
25        self.ds_offsets = np.cumsum(self.ds_lens)
datasets
ndim
ds_lens
ds_offsets