基于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
相关产品推荐
相关产品推荐

