PyTorch Lightning自定义MNIST DataLoader返回数据量异常问题求解
问题产生原因
- 对DataLoader的长度统计逻辑存在误解:
len(dataloader)返回的是批次(batch)的总数量,而非数据集的总样本数。按照配置的batch_size=256计算:- 训练集共55000样本:
ceil(55000 / 256) = 215个batch - 验证集共5000样本:
ceil(5000 / 256) = 20个batch
你观测到的215、20本身是符合预期的批次数量,并非数据缺失。
- 训练集共55000样本:
test_dataloader方法存在逻辑错误:该方法内部创建了DataLoader对象,但没有添加return语句,因此调用后返回None。- 类实例使用逻辑存在隐患:每次调用
MNISTData()都会生成一个全新的类实例,不同实例的属性相互独立。如果执行MNISTData().setup()给临时实例初始化数据后,再调用MNISTData().train_dataloader()生成新实例,新实例未执行过setup方法,会触发AttributeError异常。
修复方案
- 修正
test_dataloader方法的返回逻辑:
def test_dataloader(self): mnist_test = DataLoader(self.mnist_test, batch_size=self.batch_size) return mnist_test
- 调整类实例使用方式,复用同一个实例完成初始化和dataloader获取操作,避免重复实例化:
# 仅初始化一次数据模块实例 dm = MNISTData() dm.download() dm.setup() # 复用同一实例获取各阶段数据加载器 train_loader = dm.train_dataloader() val_loader = dm.val_dataloader() test_loader = dm.test_dataloader()
- 若需要获取数据集总样本数,访问DataLoader的dataset属性统计长度即可:
# 统计样本总数 train_sample_total = len(train_loader.dataset) # 输出55000 val_sample_total = len(val_loader.dataset) # 输出5000 test_sample_total = len(test_loader.dataset) # 输出10000 # 统计批次总数 train_batch_total = len(train_loader) # 输出215
内容的提问来源于stack exchange,提问作者이재환
相关产品推荐
相关产品推荐

