Perfectly mix datapoints from several DataFlow using their :meth:`__len__()`. Will stop when all DataFlow exhausted.
| 457 | |
| 458 | |
| 459 | class 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 | |
| 496 | class ConcatData(DataFlow): |
no outgoing calls
no test coverage detected
searching dependent graphs…