PyTorch DataLoader按index顺序无重复加载预测样本的实现问题
问题根源
设置shuffle=False未生效的核心原因有两点:
- 自定义
CompanyDataset的__getitem__方法内置随机切片逻辑,即使外部按顺序传入索引,每次取到的同索引数据也是随机片段 - 原有代码存在变量笔误:定义了
Company变量但返回时写的是station,会触发运行错误
修改方案
1. 改造CompanyDataset类,新增模式切换
新增训练/预测模式参数,预测模式下关闭随机切片,使用固定位置切片保证同索引返回结果唯一:
import os import pandas as pd import numpy as np import torch from sklearn.preprocessing import MinMaxScaler from joblib import dump from torch.utils.data import Dataset class CompanyDataset(Dataset): # 新增is_train参数,默认True为训练模式,预测时传False def __init__(self, csv_name, root_dir, training_length, forecast_window, is_train=True): """ Args: csv_file (string): Path to the csv file. root_dir (string): Directory is_train (bool): Whether it is training mode """ # load raw data file csv_file = os.path.join(root_dir, csv_name) self.df = pd.read_csv(csv_file) self.root_dir = root_dir self.transform = MinMaxScaler() self.T = training_length self.S = forecast_window self.is_train = is_train def __len__(self): # return number of sensors return len(self.df.groupby(by=["index"])) # Will pull an index between 0 and __len__. def __getitem__(self, idx): # Sensors are indexed from 1 idx = idx + 1 idx_df = self.df[self.df["index"] == idx] max_start = len(idx_df) - self.T - self.S # 模式判断:训练随机取start,预测取固定位置(这里取最后一段可用序列,可按需改成0取首段) if self.is_train: start = np.random.randint(0, max_start) else: start = max_start station = str(idx_df[["station"]][start:start + 1].values.item()) index_in = torch.tensor([i for i in range(start, start + self.T)]) index_tar = torch.tensor([i for i in range(start + self.T, start + self.T + self.S)]) _input = torch.tensor(idx_df[ ["A1","A2","A3","A4","A5","A6","A7","A8", "A9", "A10", "A11"]][ start: start + self.T].values) target = torch.tensor(idx_df[ ["A1","A2","A3","A4","A5","A6","A7","A8", "A9", "A10", "A11"]][ start + self.T: start + self.T + self.S].values) scaler = self.transform scaler.fit(_input[:, 0].unsqueeze(-1)) _input[:, 0] = torch.tensor(scaler.transform(_input[:, 0].unsqueeze(-1)).squeeze(-1)) target[:, 0] = torch.tensor(scaler.transform(target[:, 0].unsqueeze(-1)).squeeze(-1)) dump(scaler, 'scalar_item.joblib') # 修正变量名笔误,和前面定义保持一致 return index_in, index_tar, _input, target, station
2. 修改测试阶段数据加载配置
初始化测试数据集时指定预测模式,DataLoader关闭shuffle即可,无需自定义Sampler:
# 初始化测试数据集时指定is_train=False test_dataset = CompanyDataset( csv_name="你的测试csv文件名", root_dir="你的数据根目录", training_length=你的输入序列长度, forecast_window=你的预测窗口长度, is_train=False ) # shuffle设为False,默认使用SequentialSampler按顺序加载索引 test_dataloader = DataLoader(test_dataset, batch_size=1, shuffle=False)
如果你需要对每个索引做全量滑窗预测,只需要调整__len__统计所有滑窗的总数量,__getitem__里把全局id映射到对应索引和切片位置即可,当前方案已经满足你要求的1-5按顺序无重复加载的需求。
内容的提问来源于stack exchange,提问作者JalleNation
相关产品推荐
相关产品推荐

