(self)
| 306 | return data |
| 307 | |
| 308 | def __iter__(self): |
| 309 | |
| 310 | num_groups = len(self.dataset_iters) |
| 311 | ################## group weight ############ |
| 312 | total_weights = sum(self.grouped_weights) |
| 313 | assert total_weights > 0.0 |
| 314 | group_cumprobs = [sum(self.grouped_weights[:i + 1]) / total_weights |
| 315 | for i in range(len(self.grouped_weights))] |
| 316 | |
| 317 | ############################## |
| 318 | while True: |
| 319 | sequence_status = self.set_sequence_status() |
| 320 | batch_data_indexes = [] |
| 321 | added_at_least_one = False |
| 322 | |
| 323 | self.step_counter += 1 |
| 324 | step_seed = self.base_and_epoch_seed + self.step_counter |
| 325 | step_rng = random.Random(step_seed) |
| 326 | |
| 327 | random_image_num = int(step_rng.choices( |
| 328 | self.possible_nums, |
| 329 | weights=self.normalized_weights, |
| 330 | k=1 |
| 331 | )[0]) |
| 332 | random_aspect_ratio = round( |
| 333 | step_rng.uniform(self.aspect_ratio_range[0], self.aspect_ratio_range[1]), |
| 334 | 2 |
| 335 | ) |
| 336 | |
| 337 | n = step_rng.random() # Use the step_rng for reproducibility |
| 338 | group_index = 0 |
| 339 | for i, cumprob in enumerate(group_cumprobs): |
| 340 | if n < cumprob: |
| 341 | group_index = i |
| 342 | break |
| 343 | group_iter, group_name, group_dataset = self.dataset_iters[group_index] |
| 344 | |
| 345 | while True: |
| 346 | if group_name == "recon": |
| 347 | group_dataset.set_random_image_num(random_image_num) |
| 348 | group_dataset.set_random_aspect_ratio(random_aspect_ratio) |
| 349 | group_dataset.set_step_rng(step_seed) |
| 350 | sample = next(group_iter) |
| 351 | else: |
| 352 | sample = next(group_iter) |
| 353 | |
| 354 | if sample is None: |
| 355 | continue |
| 356 | # if a sample is too long, skip it |
| 357 | num_tokens = sample['num_tokens'] + 2 * len(sample['sequence_plan']) |
| 358 | |
| 359 | if num_tokens == 0 or num_tokens > self.max_num_tokens_per_sample: |
| 360 | if num_tokens == 0: |
| 361 | print(f"skip a sample with length 0") |
| 362 | else: |
| 363 | print(f"skip a sample with length {num_tokens} (exceeds max_per_sample)") |
| 364 | continue |
| 365 |
nothing calls this directly
no test coverage detected