提交 b84e4352 编写于 作者: W weishengyu

dbg

上级 6529765a
...@@ -62,7 +62,7 @@ class PKSampler(DistributedBatchSampler): ...@@ -62,7 +62,7 @@ class PKSampler(DistributedBatchSampler):
elif self.sample_method == "sample_avg_prob": elif self.sample_method == "sample_avg_prob":
counter = [] counter = []
for label_i in self.label_list: 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) self.prob_list = np.array(counter) / sum(counter)
else: else:
logger.error( logger.error(
...@@ -83,7 +83,7 @@ class PKSampler(DistributedBatchSampler): ...@@ -83,7 +83,7 @@ class PKSampler(DistributedBatchSampler):
np.random.RandomState(self.epoch).shuffle(self.label_list) np.random.RandomState(self.epoch).shuffle(self.label_list)
for i in range(len(self)): for i in range(len(self)):
batch_index = [] batch_index = []
batch_label_list = np.random.sample( batch_label_list = np.random.choice(
self.label_list, self.label_list,
size=label_per_batch, size=label_per_batch,
replace=False, replace=False,
......
Markdown is supported
0% .
You are about to add 0 people to the discussion. Proceed with caution.
先完成此消息的编辑!
想要评论请 注册