解决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
相关产品推荐
相关产品推荐

