Browse Source

fix when dataset is GeneratorDataset, sampler is DistributedSampler, dataset_size may be error

tags/v1.1.1
anzhengqi 5 years ago
parent
commit
451e488b58
1 changed files with 2 additions and 2 deletions
  1. +2
    -2
      mindspore/dataset/engine/datasets.py

+ 2
- 2
mindspore/dataset/engine/datasets.py View File

@@ -3795,10 +3795,10 @@ class GeneratorDataset(MappableDataset):
# lose attribution of '__len__' after deepcopy.
self.dataset_size = None
if hasattr(self.source, "__len__"):
if not self.num_shards:
if not isinstance(self.sampler, samplers.DistributedSampler):
self.dataset_size = len(self.source)
else:
self.dataset_size = math.ceil(len(self.source) / self.num_shards)
self.dataset_size = math.ceil(len(self.source) / self.sampler.num_shards)

rows_from_sampler = self._get_sampler_dataset_size()
if self.num_samples is not None and self.num_samples < rows_from_sampler:


Loading…
Cancel
Save