From 451e488b581ac7da79ebc027423723a488c4f33f Mon Sep 17 00:00:00 2001 From: anzhengqi Date: Wed, 27 Jan 2021 17:34:21 +0800 Subject: [PATCH] fix when dataset is GeneratorDataset, sampler is DistributedSampler, dataset_size may be error --- mindspore/dataset/engine/datasets.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/mindspore/dataset/engine/datasets.py b/mindspore/dataset/engine/datasets.py index 1a75a3b774..aee384a6fa 100644 --- a/mindspore/dataset/engine/datasets.py +++ b/mindspore/dataset/engine/datasets.py @@ -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: