(
self,
dataset: Dataset,
sampler: Sampler = None,
transform: Transform = None,
collator: Collator = None,
num_workers: int = 0,
timeout: int = 0,
preload: bool = False,
parallel_stream: bool = False,
)
| 135 | """ |
| 136 | |
| 137 | def __init__( |
| 138 | self, |
| 139 | dataset: Dataset, |
| 140 | sampler: Sampler = None, |
| 141 | transform: Transform = None, |
| 142 | collator: Collator = None, |
| 143 | num_workers: int = 0, |
| 144 | timeout: int = 0, |
| 145 | preload: bool = False, |
| 146 | parallel_stream: bool = False, |
| 147 | ): |
| 148 | if num_workers < 0: |
| 149 | raise ValueError("num_workers should not be negative") |
| 150 | |
| 151 | if timeout < 0: |
| 152 | raise ValueError("timeout should not be negative") |
| 153 | |
| 154 | self.dataset = dataset |
| 155 | self.num_workers = num_workers |
| 156 | self.timeout = timeout |
| 157 | self.preload = preload |
| 158 | self.parallel_stream = parallel_stream |
| 159 | |
| 160 | if isinstance(dataset, StreamDataset): |
| 161 | self.sampler = sampler if sampler else StreamSampler(batch_size=1) |
| 162 | assert isinstance( |
| 163 | self.sampler, StreamSampler |
| 164 | ), "types of dataset and sampler do not match" |
| 165 | if parallel_stream is False and self.num_workers > 1: |
| 166 | logger.warning( |
| 167 | "Data time will be affected by getting origin-data, please set parallel_stream in order to speed up dataloader!" |
| 168 | ) |
| 169 | self.datakind = "stream" |
| 170 | else: |
| 171 | assert isinstance( |
| 172 | dataset, Dataset |
| 173 | ), "Can not recognize this kind of dataset: %s" % type(dataset) |
| 174 | self.sampler = ( |
| 175 | sampler |
| 176 | if sampler |
| 177 | else SequentialSampler(dataset, batch_size=1, drop_last=False) |
| 178 | ) |
| 179 | assert isinstance( |
| 180 | self.sampler, MapSampler |
| 181 | ), "types of dataset and sampler do not match" |
| 182 | self.datakind = "map" |
| 183 | |
| 184 | if transform is None: |
| 185 | self.transform = PseudoTransform() |
| 186 | else: |
| 187 | self.transform = transform |
| 188 | |
| 189 | if collator is None: |
| 190 | self.collator = Collator() |
| 191 | else: |
| 192 | self.collator = collator |
| 193 | |
| 194 | if platform.system() == "Linux" and self.num_workers > 0: |
nothing calls this directly
no test coverage detected