MCPcopy Create free account
hub / github.com/MoonInTheRiver/DiffSinger / BaseDataset

Class BaseDataset

tasks/base_task.py:30–73  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

28
29
30class 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
76class BaseTask(nn.Module):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected