MCPcopy Create free account
hub / github.com/apple/ml-4m / MixtureDataset

Class MixtureDataset

fourm/data/unified_datasets.py:524–546  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

522
523
524class MixtureDataset(IterableDataset):
525 def __init__(self, data_iters, weights, modality_info):
526 self.orig_data_iters = data_iters
527 self.data_iters = [iter(data_iter) for data_iter in data_iters] # Create initial iterators
528 self.sampling_probs = np.array(weights) / sum(weights)
529 self.modality_info = modality_info
530
531 def reset_iterator(self, idx):
532 """ Reset the iterator when exhausted. """
533 self.data_iters[idx] = iter(self.orig_data_iters[idx])
534
535 def __iter__(self):
536 while True:
537 dataset_idx = np.random.choice(len(self.sampling_probs), p=self.sampling_probs)
538 try:
539 data = next(self.data_iters[dataset_idx])
540 except StopIteration: # If the iterator is exhausted
541 self.reset_iterator(dataset_idx) # Reset it
542 data = next(self.data_iters[dataset_idx])
543
544 mod_dict = make_empty_mod_dict(self.modality_info)
545 mod_dict.update(data)
546 yield mod_dict
547
548
549def build_mixture_dataloader(data_iters, weights, modality_info, batch_size, num_workers, epoch_size, num_gpus):

Callers 1

build_mixture_dataloaderFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected