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)