| 60 | from typing import Any, Dict, List, Set |
| 61 | |
| 62 | class JointLoader(torchdata.IterableDataset): |
| 63 | def __init__(self, loaders, key_dataset): |
| 64 | dataset_names = [] |
| 65 | for key, loader in loaders.items(): |
| 66 | name = "{}".format(key.split('_')[0]) |
| 67 | setattr(self, name, loader) |
| 68 | dataset_names += [name] |
| 69 | self.dataset_names = dataset_names |
| 70 | self.key_dataset = key_dataset |
| 71 | |
| 72 | def __iter__(self): |
| 73 | for batch in zip(*[getattr(self, name) for name in self.dataset_names]): |
| 74 | yield {key: batch[i] for i, key in enumerate(self.dataset_names)} |
| 75 | |
| 76 | def __len__(self): |
| 77 | return len(getattr(self, self.key_dataset)) |
| 78 | |
| 79 | def filter_images_with_only_crowd_annotations(dataset_dicts, dataset_names): |
| 80 | """ |
no outgoing calls
no test coverage detected