diff --git a/gigashuffle/multiprocess.py b/gigashuffle/multiprocess.py index 193e535..e158d3e 100644 --- a/gigashuffle/multiprocess.py +++ b/gigashuffle/multiprocess.py @@ -407,10 +407,16 @@ def initialize_reader(config: DataloaderConfig, proc_idx: int, queue_name: str) def copy_to_reader_buffer(reader_buffer: Buffer, shuffle_buffer: Buffer, idx_list: list[int]) -> None: + idx = torch.as_tensor(idx_list, dtype=torch.int64) for buffer_idx in range(len(shuffle_buffer)): - for k in shuffle_buffer[buffer_idx].keys(): - reader_buffer[buffer_idx][k][:] = shuffle_buffer[buffer_idx][k][idx_list] - reader_buffer[0][INDEX_KEY].copy_(torch.as_tensor(idx_list)) + for key in shuffle_buffer[buffer_idx]: + torch.index_select( + shuffle_buffer[buffer_idx][key], + 0, + idx, + out=reader_buffer[buffer_idx][key], + ) + reader_buffer[0][INDEX_KEY].copy_(idx) def send_reader_buffer(ready_q: SimpleQueue[tuple[Buffer, int]], ready_e: Event, reader_buffer: Buffer, proc_idx: int) -> None: