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

解决PyTorch IterableDataset多Worker读取大文件时产生重复样本的问题

这个问题我之前也踩过坑!核心原因很简单:每个DataLoader Worker都会独立创建一个Dataset实例——当你设置num_workers=2时,会有两个CustomIterableDatasetv1对象分别在两个子进程里运行,每个都从头打开文件并读取完整内容,自然就会输出重复样本了。

要解决这个问题,我们需要让每个Worker只负责读取文件的一部分内容,避免重复处理。PyTorch提供了torch.utils.data.get_worker_info()工具,可以帮我们获取当前Worker的身份信息,从而实现数据的分片处理。下面给你两种可行的解决方案:

方案一:基于行数的分片(适合中小文件)

这种方式先计算文件的总行数,然后给每个Worker分配对应的行区间,实现起来简单直观:

from torch.utils.data import IterableDataset, DataLoader
import torch.utils.data

class CustomIterableDatasetv2(IterableDataset):
    def __init__(self, filename):
        self.filename = filename

    def preprocess(self, text):
        text_pp = text.lower().strip()
        return text_pp

    def line_mapper(self, line):
        text, label = line.split('-')
        text = self.preprocess(text)
        return text, label

    def __iter__(self):
        worker_info = torch.utils.data.get_worker_info()
        if worker_info is None:
            # 单Worker模式:直接读取全部内容
            file_itr = open(self.filename)
        else:
            # 多Worker模式:计算当前Worker的处理范围
            num_workers = worker_info.num_workers
            worker_id = worker_info.id
            
            # 先统计文件总行数(如果预先知道行数,可以直接传入优化性能)
            with open(self.filename) as f:
                total_lines = sum(1 for _ in f)
            
            # 分配行区间:最后一个Worker处理剩余所有行
            lines_per_worker = total_lines // num_workers
            start_line = worker_id * lines_per_worker
            end_line = start_line + lines_per_worker
            if worker_id == num_workers - 1:
                end_line = total_lines
            
            # 跳过不属于当前Worker的行,生成迭代器
            file = open(self.filename)
            # 跳过前start_line行
            for _ in range(start_line):
                next(file)
            # 读取当前Worker负责的行
            file_itr = []
            for _ in range(end_line - start_line):
                try:
                    line = next(file)
                    file_itr.append(line)
                except StopIteration:
                    break
            file.close()
        
        # 映射处理每一行
        mapped_itr = map(self.line_mapper, file_itr)
        return mapped_itr

测试验证

运行以下代码:

base_dataset = CustomIterableDatasetv2("testfile.txt")
dataloader = DataLoader(base_dataset, batch_size=1, num_workers=2)
for X, y in dataloader:
    print(X, y)

你会得到无重复的输出(顺序可能因为Worker并行略有变化,但所有样本只会出现一次)。

方案二:基于字节位置的分片(适合超大文件)

如果你的文件大到统计行数都很耗时,可以用字节位置来分片,同时保证不拆分完整的行:

class CustomIterableDatasetv3(IterableDataset):
    def __init__(self, filename):
        self.filename = filename

    def preprocess(self, text):
        text_pp = text.lower().strip()
        return text_pp

    def line_mapper(self, line):
        text, label = line.split('-')
        text = self.preprocess(text)
        return text, label

    def __iter__(self):
        worker_info = torch.utils.data.get_worker_info()
        with open(self.filename, 'rb') as file:
            # 获取文件总字节数
            file.seek(0, 2)
            total_size = file.tell()
            file.seek(0)

            if worker_info is None:
                # 单Worker:读取全部内容并按行拆分
                content = file.read().decode()
                file_itr = content.splitlines()
            else:
                num_workers = worker_info.num_workers
                worker_id = worker_info.id
                
                # 计算每个Worker的字节区间
                chunk_size = total_size // num_workers
                start = worker_id * chunk_size
                end = start + chunk_size if worker_id != num_workers -1 else total_size

                # 调整起始位置:跳到下一行开头,避免拆分行
                file.seek(start)
                if start != 0:
                    file.readline()
                    start = file.tell()
                
                # 调整结束位置:回到上一行末尾,避免拆分行
                file.seek(end)
                if end != total_size:
                    while file.tell() > start:
                        file.seek(-1, 1)
                        if file.read(1) == b'\n':
                            break
                    end = file.tell()
                
                # 读取当前Worker负责的内容并按行拆分
                file.seek(start)
                content = file.read(end - start).decode()
                file_itr = content.splitlines()
        
        mapped_itr = map(self.line_mapper, file_itr)
        return mapped_itr

这种方式不需要预先统计行数,直接通过字节位置分片,更适合处理GB级以上的超大文件。

关键要点总结

  • 利用get_worker_info()区分单/多Worker环境,获取当前Worker的id和总num_workers。
  • 多Worker模式下,必须给每个Worker分配独立的处理范围,避免重复读取。
  • 两种分片方式按需选择:行数分片简单,字节分片高效适合超大文件。

内容的提问来源于stack exchange,提问作者eljiwo

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.30 21:27:41