From b84e4352b11104176b3baa9c594c404ca3b07c41 Mon Sep 17 00:00:00 2001 From: weishengyu Date: Sun, 26 Sep 2021 14:28:12 +0800 Subject: [PATCH] dbg --- ppcls/data/dataloader/pk_sampler.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/ppcls/data/dataloader/pk_sampler.py b/ppcls/data/dataloader/pk_sampler.py index 4e872ac0..ef65be94 100644 --- a/ppcls/data/dataloader/pk_sampler.py +++ b/ppcls/data/dataloader/pk_sampler.py @@ -62,7 +62,7 @@ class PKSampler(DistributedBatchSampler): elif self.sample_method == "sample_avg_prob": counter = [] for label_i in self.label_list: - counter.append(len(self.label_list[label_i])) + counter.append(len(self.label_dict[label_i])) self.prob_list = np.array(counter) / sum(counter) else: logger.error( @@ -83,7 +83,7 @@ class PKSampler(DistributedBatchSampler): np.random.RandomState(self.epoch).shuffle(self.label_list) for i in range(len(self)): batch_index = [] - batch_label_list = np.random.sample( + batch_label_list = np.random.choice( self.label_list, size=label_per_batch, replace=False, -- GitLab