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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 16:39:01