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

Function split_to_even_chunks

train/monkey_patch.py:78–97  ·  view source on GitHub ↗

Split a list of indices into `chunks` chunks of roughly equal lengths.

(indices, lengths, num_chunks)

Source from the content-addressed store, hash-verified

76
77# copy from https://github.com/haotian-liu/LLaVA/blob/main/llava/train/llava_trainer.py#L38
78def split_to_even_chunks(indices, lengths, num_chunks):
79 """
80 Split a list of indices into `chunks` chunks of roughly equal lengths.
81 """
82
83 if len(indices) % num_chunks != 0:
84 return [indices[i::num_chunks] for i in range(num_chunks)]
85
86 num_indices_per_chunk = len(indices) // num_chunks
87
88 chunks = [[] for _ in range(num_chunks)]
89 chunks_lengths = [0 for _ in range(num_chunks)]
90 for index in indices:
91 shortest_chunk = chunks_lengths.index(min(chunks_lengths))
92 chunks[shortest_chunk].append(index)
93 chunks_lengths[shortest_chunk] += lengths[index]
94 if len(chunks[shortest_chunk]) == num_indices_per_chunk:
95 chunks_lengths[shortest_chunk] = float('inf')
96
97 return chunks
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):

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected