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

多CSV文件滑动窗口PyTorch Dataset报错排查与实现求助

问题分析与修复方案

错误原因拆解

  1. __len__方法逻辑错误:当前返回的是第一个文件名的字符长度,完全不符合数据集样本总数的定义,这也是触发TypeError的直接原因之一。
  2. 实例方法调用错误:__getitem__中直接调用read_file(filename),而非self.read_file(filename),会导致无法找到实例方法。
  3. 滑动窗口生成逻辑错误:原代码生成的窗口包含超出有效范围的片段(部分窗口长度不足32),且未按需求只保留前10个完整窗口。
  4. Dataset样本粒度错误:原代码把整个文件的所有窗口作为一个样本返回,不符合PyTorch Dataset「每个样本对应一个数据单元」的设计原则,应将单个滑动窗口作为独立样本。

修正后的完整代码

import os
import pandas as pd
from torch.utils.data import Dataset

class CustomDataset(Dataset):
    def __init__(self, data_folder, window_size):
        self.data_folder = data_folder
        # 仅读取CSV文件,过滤无关文件
        self.data_file_list = [
            os.path.join(data_folder, file) 
            for file in os.listdir(data_folder) 
            if file.endswith('.csv')
        ]
        self.window_size = window_size
        # 每个文件固定生成10个窗口(符合需求)
        self.per_file_windows = 10
        # 总样本数 = 文件数量 × 单文件窗口数
        self.total_samples = len(self.data_file_list) * self.per_file_windows

    def __len__(self):
        return self.total_samples

    def __getitem__(self, idx):
        # 计算当前样本所属文件索引与文件内窗口索引
        file_idx = idx // self.per_file_windows
        window_idx = idx % self.per_file_windows
        
        filename = self.data_file_list[file_idx]
        features, labels = self.read_file(filename)
        
        # 提取对应窗口的特征与标签
        start = window_idx * self.window_size
        end = start + self.window_size
        x = features[start:end].values
        y = labels.iloc[end-1]  # 默认取窗口最后一个数据点的标签,可按需调整
        
        return x, y

    def read_file(self, filename):
        data = pd.read_csv(filename)
        # 移除无关列
        data = data.drop(["file_name", "class_name"], axis=1)
        features = data.drop(["class_no"], axis=1)
        labels = data["class_no"]
        
        # 保留前320个数据点(对应10个窗口),丢弃最后10个无效点
        features = features.iloc[:320]
        labels = labels.iloc[:320]
        
        return features, labels

关键修正说明

  • 样本粒度调整:将单个滑动窗口作为Dataset的一个样本,__len__返回所有文件的总窗口数,契合PyTorch数据加载逻辑。
  • 窗口生成逻辑:通过window_idx * window_size计算窗口起始位置,确保每个窗口都是连续的32个数据点,严格保留前10个有效窗口。
  • 文件路径处理:使用os.path.join拼接完整文件路径,避免因工作目录问题导致文件查找失败。
  • 标签处理:默认取窗口最后一个数据点的标签作为该窗口的标签,可根据实际需求修改(如取窗口内多数标签、平均值等)。
  • 文件过滤:仅读取.csv后缀的文件,避免目录中其他无关文件干扰。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 00:55:27