You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何让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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.16 13:07:43