MCPcopy Create free account
hub / github.com/tensorpack/tensorpack / RandomMixData

Class RandomMixData

tensorpack/dataflow/common.py:459–493  ·  view source on GitHub ↗

Perfectly mix datapoints from several DataFlow using their :meth:`__len__()`. Will stop when all DataFlow exhausted.

Source from the content-addressed store, hash-verified

457
458
459class RandomMixData(RNGDataFlow):
460 """
461 Perfectly mix datapoints from several DataFlow using their
462 :meth:`__len__()`. Will stop when all DataFlow exhausted.
463 """
464
465 def __init__(self, df_lists):
466 """
467 Args:
468 df_lists (list): a list of DataFlow.
469 All DataFlow must implement ``__len__()``.
470 """
471 super(RandomMixData, self).__init__()
472 self.df_lists = df_lists
473 self.sizes = [len(k) for k in self.df_lists]
474
475 def reset_state(self):
476 super(RandomMixData, self).reset_state()
477 for d in self.df_lists:
478 d.reset_state()
479
480 def __len__(self):
481 return sum(self.sizes)
482
483 def __iter__(self):
484 sums = np.cumsum(self.sizes)
485 idxs = np.arange(self.__len__())
486 self.rng.shuffle(idxs)
487 idxs = np.array(list(map(
488 lambda x: np.searchsorted(sums, x, 'right'), idxs)))
489 itrs = [k.__iter__() for k in self.df_lists]
490 assert idxs.max() == len(itrs) - 1, "{}!={}".format(idxs.max(), len(itrs) - 1)
491 for k in idxs:
492 yield next(itrs[k])
493 # TODO run till exception
494
495
496class ConcatData(DataFlow):

Callers 2

get_dataFunction · 0.85
get_configFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…