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

基于8个DataFrame创建PyTorch DataLoader遇批量异常问题求助

问题分析与解决

核心错误原因

你的Dataset类中,self.num_samples = len(data_frames_list[0])这一行存在逻辑错误:Pandas DataFrame的len()返回的是行数(这里是11),但你的样本总数是DataFrame的列数(4,232,460)。这就导致Dataset的总长度被错误设置为11,所以DataLoader用batch_size=8时,只能生成1个批次(11//8=1),完全不符合预期。

修正后的Dataset实现

方案1:直接修正样本数计算(简单易实现)

仅需将样本数的计算改为取DataFrame的列数,即df.shape[1]:

import torch
from torch.utils.data import Dataset

class CustomDataset(Dataset):
    def __init__(self, data_frames_list, target_list):
        self.data_frames_list = data_frames_list
        self.target_list = target_list
        # 样本数为DataFrame的列数,而非行数
        self.num_samples = data_frames_list[0].shape[1]

    def __len__(self):
        return self.num_samples

    def __getitem__(self, idx):
        samples = [df.iloc[:, idx].values for df in self.data_frames_list]
        samples = torch.FloatTensor(samples)  # 形状: (8, 11)
        target = torch.FloatTensor([self.target_list[idx]])
        return samples, target

方案2:提前预处理为张量(加载效率更高)

如果内存允许,建议提前将所有DataFrame合并为一个大张量,避免在__getitem__中频繁调用Pandas的iloc,大幅提升数据加载速度:

import torch
import numpy as np
from torch.utils.data import Dataset

class CustomDataset(Dataset):
    def __init__(self, data_frames_list, target_list):
        # 将所有DataFrame按特征维度堆叠,得到形状(8, 11, 4232460)的数组
        # 转置为(4232460, 8, 11),方便按样本索引直接取值
        self.data = torch.FloatTensor(
            np.stack([df.values for df in data_frames_list], axis=0)
        ).permute(2, 0, 1)
        self.targets = torch.FloatTensor(target_list)
        self.num_samples = self.data.shape[0]

    def __len__(self):
        return self.num_samples

    def __getitem__(self, idx):
        # 直接按索引获取样本,形状(8,11)
        sample = self.data[idx]
        target = self.targets[idx]
        return sample, target

验证效果

使用修正后的Dataset创建DataLoader,就能正常生成对应数量的批次:

from torch.utils.data import DataLoader

# 假设你已经有了train_dfs(8个DataFrame)和train_targets(长度4232460的标签列表)
train_dataset = CustomDataset(train_dfs, train_targets)

# 设置任意合理的batch_size,比如64
batch_size = 64
train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, drop_last=True)

print(f"样本总数: {len(train_dataset)}")  # 输出4232460
print(f"批次数量: {len(train_loader)}")  # 输出4232460//64=66132(如果drop_last=True)

# 查看单个批次的形状
for batch_data, batch_target in train_loader:
    print(f"批次数据形状: {batch_data.shape}")  # (64, 8, 11)
    print(f"批次标签形状: {batch_target.shape}")  # (64, 1)
    break

注意事项

  • 确保target_list的长度和DataFrame的列数完全一致(都是4,232,460),否则会出现索引越界错误。
  • 如果内存不足以一次性加载所有数据,优先使用方案1,或者考虑分块加载数据。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 08:28:41