From 263b8fc916665711743fb0c1611f85db44b4a021 Mon Sep 17 00:00:00 2001 From: YassineYousfi Date: Mon, 27 Jul 2026 17:54:29 -0700 Subject: [PATCH] use index select --- gigashuffle/multiprocess.py | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) 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: