PyTorch自定义DataLoader读取多份大型CSV文件的跨文件索引错误解决
自定义多CSV滑动窗口数据集修复方案
问题根因
- 样本总数计算错误:原代码直接将所有CSV的行数总和作为数据集长度,未考虑滑动窗口的前置数据要求。每个CSV仅能独立产出
行数 - 历史窗口大小个有效样本,不允许跨文件使用其他CSV的行作为历史数据。 - 索引映射逻辑错误:原代码将所有CSV的行视为全局连续序列,未做文件间的样本隔离,切换到第二个CSV时会触发错误的索引偏移逻辑,导致读取重复的行数据。
- 文件读取效率低:原代码遍历整个CSV文件获取目标行,读取到目标行范围后可直接终止,降低IO开销。
修复后完整代码
import numpy as np from functools import lru_cache from pathlib import Path from pprint import pprint from torch.utils.data import Dataset, DataLoader import torch @lru_cache() def get_sample_count_by_file(path: Path) -> int: c = 0 with path.open() as f: for line in f: c += 1 return c class CSVDataset(Dataset): def __init__(self, csv_directory: str, extension: str = ".csv", history_window: int = 2): self.directory = Path(csv_directory) self.history_window = history_window # 存储每个文件的路径、总行数、有效样本数、全局样本起始索引 self.files = [] total_sample = 0 for f in sorted(self.directory.iterdir()): if f.suffix == extension: total_line = get_sample_count_by_file(f) # 有效样本数 = 总行数 - 窗口大小,行数不足的文件自动跳过 valid_sample = max(0, total_line - self.history_window) if valid_sample > 0: self.files.append((f, total_line, valid_sample, total_sample)) total_sample += valid_sample self._sample_count = total_sample def __len__(self): return self._sample_count def __getitem__(self, idx): # 定位当前索引所属的文件 target_file = None sample_offset_in_file = 0 for file_, total_line, valid_sample, start_idx in self.files: if start_idx <= idx < start_idx + valid_sample: target_file = file_ sample_offset_in_file = idx - start_idx break if not target_file: raise IndexError("Index out of range") # 计算当前行在文件内的索引,以及需要读取的行范围 current_line_idx = self.history_window + sample_offset_in_file start_read = current_line_idx - self.history_window end_read = current_line_idx data = [] with target_file.open() as f: for i, line in enumerate(f): # 读到目标行范围外直接终止,减少无效IO if i > end_read: break if start_read <= i <= end_read: for v in line.strip().split(","): data.append(float(v.strip())) data = np.array(data) return torch.from_numpy(data) dataset = CSVDataset("<PATH CONTAINING CSVs>") loader = DataLoader(dataset, batch_size=1) pprint(list(enumerate(loader)))
关键修改说明
- 新增文件级样本隔离逻辑:每个CSV的有效样本独立计算,不会跨文件复用历史数据,完全符合业务预期。
- 修正索引映射逻辑:通过预存每个文件的有效样本起始/结束区间,直接定位到样本所属的文件以及文件内的偏移,移除了原代码中错误的索引修正逻辑。
- 优化文件读取逻辑:读取到目标行范围后直接终止循环,大幅降低大文件的读取开销。
- 兼容PyTorch官方规范:显式继承
Dataset类,避免潜在的DataLoader兼容问题。
内容的提问来源于stack exchange,提问作者Anto
相关产品推荐
相关产品推荐

