多CSV文件滑动窗口PyTorch Dataset报错排查与实现求助
问题分析与修复方案
错误原因拆解
__len__方法逻辑错误:当前返回的是第一个文件名的字符长度,完全不符合数据集样本总数的定义,这也是触发TypeError的直接原因之一。- 实例方法调用错误:
__getitem__中直接调用read_file(filename),而非self.read_file(filename),会导致无法找到实例方法。 - 滑动窗口生成逻辑错误:原代码生成的窗口包含超出有效范围的片段(部分窗口长度不足32),且未按需求只保留前10个完整窗口。
- 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
相关产品推荐
相关产品推荐

