(lengths, batch_size, world_size, generator=None, merge=True)
| 98 | |
| 99 | # copy from https://github.com/haotian-liu/LLaVA/blob/main/llava/train/llava_trainer.py#L88 |
| 100 | def get_length_grouped_indices(lengths, batch_size, world_size, generator=None, merge=True): |
| 101 | # We need to use torch for the random part as a distributed sampler will set the random seed for torch. |
| 102 | indices = torch.randperm(len(lengths), generator=generator) |
| 103 | megabatch_size = world_size * batch_size |
| 104 | megabatches = [indices[i : i + megabatch_size].tolist() for i in range(0, len(lengths), megabatch_size)] |
| 105 | megabatches = [sorted(megabatch, key=lambda i: lengths[i], reverse=True) for megabatch in megabatches] |
| 106 | megabatches = [split_to_even_chunks(megabatch, lengths, world_size) for megabatch in megabatches] |
| 107 | |
| 108 | return [i for megabatch in megabatches for batch in megabatch for i in batch] |
| 109 | |
| 110 | # modified from https://github.com/haotian-liu/LLaVA/blob/main/llava/train/llava_trainer.py#L99 |
| 111 | class LengthGroupedSampler(Sampler): |
no test coverage detected