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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 19:39:03