Dataset wrapper that randomly subsamples the original dataset. Args: original_dataset (torch.utils.data.Dataset): The original dataset to be subsampled. dataset_size (int): The size of the subsampled dataset. seed (int): The seed to use for selec
(self, original_dataset, dataset_size, seed=0, return_orig_idx=False)
| 49 | |
| 50 | class SubsampleDatasetWrapper(Dataset): |
| 51 | def __init__(self, original_dataset, dataset_size, seed=0, return_orig_idx=False): |
| 52 | """ |
| 53 | Dataset wrapper that randomly subsamples the original dataset. |
| 54 | |
| 55 | Args: |
| 56 | original_dataset (torch.utils.data.Dataset): The original dataset to be subsampled. |
| 57 | dataset_size (int): The size of the subsampled dataset. |
| 58 | seed (int): The seed to use for selecting the subset of indices of the original dataset. |
| 59 | return_orig_idx (bool): Whether to return the original index of the item in the original dataset. |
| 60 | """ |
| 61 | self.original_dataset = original_dataset |
| 62 | self.dataset_size = dataset_size or len(original_dataset) |
| 63 | self.return_orig_idx = return_orig_idx |
| 64 | np.random.seed(seed) |
| 65 | self.indices = np.random.permutation(len(self.original_dataset))[:self.dataset_size] |
| 66 | |
| 67 | def __getitem__(self, index): |
| 68 | """ |
nothing calls this directly
no outgoing calls
no test coverage detected