| 58 | |
| 59 | |
| 60 | class VariableVideoBatchSampler(DistributedSampler): |
| 61 | |
| 62 | def __init__( |
| 63 | self, |
| 64 | dataset, |
| 65 | bucket_config: dict, |
| 66 | num_replicas: Optional[int] = None, |
| 67 | rank: Optional[int] = None, |
| 68 | shuffle: bool = True, |
| 69 | seed: int = 0, |
| 70 | drop_last: bool = False, |
| 71 | verbose: bool = False, |
| 72 | num_bucket_build_workers: int = 1, |
| 73 | ) -> None: |
| 74 | super().__init__( |
| 75 | dataset=dataset, |
| 76 | num_replicas=num_replicas, |
| 77 | rank=rank, |
| 78 | shuffle=shuffle, |
| 79 | seed=seed, |
| 80 | drop_last=drop_last, |
| 81 | ) |
| 82 | self.dataset = dataset |
| 83 | self.bucket = Bucket(bucket_config) |
| 84 | self.verbose = verbose |
| 85 | self.last_micro_batch_access_index = 0 |
| 86 | self.approximate_num_batch = None |
| 87 | |
| 88 | self._get_num_batch_cached_bucket_sample_dict = None |
| 89 | self.num_bucket_build_workers = num_bucket_build_workers |
| 90 | |
| 91 | def __iter__(self) -> Iterator[List[int]]: |
| 92 | if self._get_num_batch_cached_bucket_sample_dict is not None: |
| 93 | bucket_sample_dict = self._get_num_batch_cached_bucket_sample_dict |
| 94 | self._get_num_batch_cached_bucket_sample_dict = None |
| 95 | else: |
| 96 | bucket_sample_dict = self.group_by_bucket() |
| 97 | if self.verbose: |
| 98 | self._print_bucket_info(bucket_sample_dict) |
| 99 | |
| 100 | g = torch.Generator() |
| 101 | g.manual_seed(self.seed + self.epoch) |
| 102 | bucket_micro_batch_count = OrderedDict() |
| 103 | bucket_last_consumed = OrderedDict() |
| 104 | |
| 105 | # process the samples |
| 106 | for bucket_id, data_list in bucket_sample_dict.items(): |
| 107 | # handle droplast |
| 108 | bs_per_gpu = self.bucket.get_batch_size(bucket_id) |
| 109 | remainder = len(data_list) % bs_per_gpu |
| 110 | |
| 111 | if remainder > 0: |
| 112 | if not self.drop_last: |
| 113 | # if there is remainder, we pad to make it divisible |
| 114 | data_list += data_list[:bs_per_gpu - remainder] |
| 115 | else: |
| 116 | # we just drop the remainder to make it divisible |
| 117 | data_list = data_list[:-remainder] |
nothing calls this directly
no outgoing calls
no test coverage detected