如何让PyTorch多进程DataLoader按FIFO返回已加载完成的批次?
问题:PyTorch DataLoader 优先返回已加载完成的批次
我用PyTorch实现了一个IterableDataset加载数据,但部分样本加载耗时较长,其余样本加载速度很快。当前DataLoader会按Worker顺序轮流返回批次,我希望它能在某个批次的所有样本加载完成后立即返回该批次,无需等待Worker轮次。
示例代码
import torch import math import time class MyIterableDataset(torch.utils.data.IterableDataset): def __init__(self, start, end): super(MyIterableDataset).__init__() assert end > start, "this example code only works with end >= start" self.start = start self.end = end def give_data(self, start, end): for i in range(start, end): if i > 10: time.sleep(2) yield i def __iter__(self): worker_info = torch.utils.data.get_worker_info() if worker_info is None: # 单进程加载,返回完整迭代器 iter_start = self.start iter_end = self.end else: # 多进程Worker环境 # 拆分工作负载 per_worker = int(math.ceil((self.end - self.start) / float(worker_info.num_workers))) worker_id = worker_info.id iter_start = self.start + worker_id * per_worker iter_end = min(iter_start + per_worker, self.end) return self.give_data(iter_start, iter_end) if __name__ == "__main__": ds = MyIterableDataset(start=0, end=20) # 双进程加载 for item in (torch.utils.data.DataLoader(ds, num_workers=2, batch_size=2)): print(item)
当前输出
tensor([0, 1]) # 快速加载完成 tensor([10, 11]) # 加载缓慢 tensor([2, 3]) # 快速加载完成 tensor([12, 13]) # 加载缓慢 tensor([4, 5]) # 快速加载完成 tensor([14, 15]) # 加载缓慢 tensor([6, 7]) # 快速加载完成 tensor([16, 17]) # 加载缓慢 tensor([8, 9]) # 快速加载完成 tensor([18, 19]) # 加载缓慢
期望输出(优先返回已加载完成的快批次)
tensor([0, 1]) # 快速加载完成 tensor([2, 3]) # 快速加载完成 tensor([4, 5]) # 快速加载完成 tensor([6, 7]) # 快速加载完成 tensor([10, 11]) # 加载缓慢 tensor([8, 9]) # 快速加载完成 tensor([12, 13]) # 加载缓慢 tensor([14, 15]) # 加载缓慢 tensor([16, 17]) # 加载缓慢 tensor([18, 19]) # 加载缓慢
解决方案
默认PyTorch DataLoader在多进程模式下会按Worker顺序依次获取批次,导致慢Worker的批次阻塞快Worker的后续输出。要实现“先完成先返回”的效果,可通过以下方式解决:
方法1:调整Worker任务分配(缓解阻塞)
默认的连续区间拆分会让一个Worker全处理慢样本,另一个全处理快样本。改为按步长分配任务,让每个Worker混合快慢样本,减少阻塞:
def __iter__(self): worker_info = torch.utils.data.get_worker_info() if worker_info is None: iter_start = self.start iter_end = self.end indices = range(iter_start, iter_end) else: num_workers = worker_info.num_workers worker_id = worker_info.id # 按步长分配,每个Worker取间隔num_workers的样本 indices = range(self.start + worker_id, self.end, num_workers) return self.give_data_from_indices(indices) # 新增方法:从索引列表生成数据 def give_data_from_indices(self, indices): for i in indices: if i > 10: time.sleep(2) yield i
方法2:自定义异步DataLoader(完全实现需求)
让Worker异步加载批次,用队列收集已完成的批次,主线程从队列取结果,不受Worker顺序限制:
import torch import math import time import torch.multiprocessing as mp class AsyncIterableDataset(torch.utils.data.IterableDataset): def __init__(self, start, end): super().__init__() assert end > start self.start = start self.end = end def give_data(self, indices): for i in indices: if i > 10: time.sleep(2) yield i def worker_task(worker_id, indices, batch_size, queue): batch = [] for idx in indices: batch.append(idx) if len(batch) == batch_size: queue.put(torch.tensor(batch)) batch = [] if batch: queue.put(torch.tensor(batch)) queue.put(None) # 标记Worker任务完成 def async_dataloader(dataset, num_workers=2, batch_size=2): total_indices = list(range(dataset.start, dataset.end)) per_worker = math.ceil(len(total_indices)/num_workers) worker_indices = [total_indices[i*per_worker:(i+1)*per_worker] for i in range(num_workers)] queue = mp.Queue() processes = [] for i in range(num_workers): p = mp.Process(target=worker_task, args=(i, worker_indices[i], batch_size, queue)) p.start() processes.append(p) finished_workers = 0 while finished_workers < num_workers: item = queue.get() if item is None: finished_workers +=1 else: yield item for p in processes: p.join() if __name__ == "__main__": ds = AsyncIterableDataset(start=0, end=20) for item in async_dataloader(ds, num_workers=2, batch_size=2): print(item)
方法3:用concurrent.futures实现异步加载(简化版)
将所有批次加载任务提交到进程池,按完成顺序返回结果:
import torch import math import time import concurrent.futures class MyIterableDataset(torch.utils.data.IterableDataset): def __init__(self, start, end): super().__init__() assert end > start self.start = start self.end = end def load_batch(self, batch_indices): batch = [] for i in batch_indices: if i >10: time.sleep(2) batch.append(i) return torch.tensor(batch) def get_batches(start, end, batch_size): indices = list(range(start, end)) for i in range(0, len(indices), batch_size): yield indices[i:i+batch_size] if __name__ == "__main__": ds = MyIterableDataset(0,20) batches = list(get_batches(0,20,2)) with concurrent.futures.ProcessPoolExecutor(max_workers=2) as executor: futures = [executor.submit(ds.load_batch, batch) for batch in batches] # 按完成顺序获取结果 for future in concurrent.futures.as_completed(futures): print(future.result())
内容的提问来源于stack exchange,提问作者tomotom12
相关产品推荐
相关产品推荐

