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

PyTorch Lightning自定义MNIST DataLoader返回数据量异常问题求解

问题产生原因
  1. 对DataLoader的长度统计逻辑存在误解:len(dataloader)返回的是批次(batch)的总数量,而非数据集的总样本数。按照配置的batch_size=256计算:
    • 训练集共55000样本:ceil(55000 / 256) = 215个batch
    • 验证集共5000样本:ceil(5000 / 256) = 20个batch
      你观测到的215、20本身是符合预期的批次数量,并非数据缺失。
  2. test_dataloader方法存在逻辑错误:该方法内部创建了DataLoader对象,但没有添加return语句,因此调用后返回None。
  3. 类实例使用逻辑存在隐患:每次调用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,提问作者이재환

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 14:39:03