| 28 | |
| 29 | |
| 30 | class BaseDataset(torch.utils.data.Dataset): |
| 31 | def __init__(self, shuffle): |
| 32 | super().__init__() |
| 33 | self.hparams = hparams |
| 34 | self.shuffle = shuffle |
| 35 | self.sort_by_len = hparams['sort_by_len'] |
| 36 | self.sizes = None |
| 37 | |
| 38 | @property |
| 39 | def _sizes(self): |
| 40 | return self.sizes |
| 41 | |
| 42 | def __getitem__(self, index): |
| 43 | raise NotImplementedError |
| 44 | |
| 45 | def collater(self, samples): |
| 46 | raise NotImplementedError |
| 47 | |
| 48 | def __len__(self): |
| 49 | return len(self._sizes) |
| 50 | |
| 51 | def num_tokens(self, index): |
| 52 | return self.size(index) |
| 53 | |
| 54 | def size(self, index): |
| 55 | """Return an example's size as a float or tuple. This value is used when |
| 56 | filtering a dataset with ``--max-positions``.""" |
| 57 | size = min(self._sizes[index], hparams['max_frames']) |
| 58 | return size |
| 59 | |
| 60 | def ordered_indices(self): |
| 61 | """Return an ordered list of indices. Batches will be constructed based |
| 62 | on this order.""" |
| 63 | if self.shuffle: |
| 64 | indices = np.random.permutation(len(self)) |
| 65 | if self.sort_by_len: |
| 66 | indices = indices[np.argsort(np.array(self._sizes)[indices], kind='mergesort')] |
| 67 | else: |
| 68 | indices = np.arange(len(self)) |
| 69 | return indices |
| 70 | |
| 71 | @property |
| 72 | def num_workers(self): |
| 73 | return int(os.getenv('NUM_WORKERS', hparams['ds_workers'])) |
| 74 | |
| 75 | |
| 76 | class BaseTask(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected