MCPcopy Create free account
hub / github.com/SpatialVLA/SpatialVLA / get_length_grouped_indices

Function get_length_grouped_indices

train/monkey_patch.py:100–108  ·  view source on GitHub ↗
(lengths, batch_size, world_size, generator=None, merge=True)

Source from the content-addressed store, hash-verified

98
99# copy from https://github.com/haotian-liu/LLaVA/blob/main/llava/train/llava_trainer.py#L88
100def 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
111class LengthGroupedSampler(Sampler):

Callers 1

__iter__Method · 0.85

Calls 1

split_to_even_chunksFunction · 0.85

Tested by

no test coverage detected