PyTorch Lightning中CIFAR10DataModule出现train_loader相关报错如何解决
错误原因及解决方案
错误1:AttributeError: 'CIFAR10DataModule' object has no attribute 'train_loader'
- 根本原因:PyTorch Lightning Bolts的
CIFAR10DataModule类原生没有名为train_loader的属性,你参考的示例大概率是旧版本PL Bolts的写法,或者是开发者二次封装后的自定义属性名,新版本官方API已经调整了相关命名。
错误2:TypeError: 'method' object is not iterable
- 根本原因:
train_dataloader是类的可调用方法,不是直接返回数据加载器的属性,你没有加()调用方法,直接遍历方法对象肯定会触发类型报错。除此之外你还漏掉了LightningDataModule必须的初始化步骤:数据准备和数据集拆分配置。
完整可运行的修复代码
import torch from pl_bolts.datamodules import CIFAR10DataModule from torch.optim import Adam from torch.nn.functional import cross_entropy # 初始化数据模块,可自定义batch_size、数据存储路径、num_workers等参数 dm = CIFAR10DataModule(batch_size=32, data_dir='./') # 执行数据下载、格式预处理步骤 dm.prepare_data() # 执行数据集拆分(训练/验证/测试集划分),stage参数fit对应训练验证阶段 dm.setup(stage='fit') # 此处为示例占位,替换为你自己的backbone和finetune_layer定义即可 backbone = torch.nn.Identity() finetune_layer = torch.nn.Linear(3*32*32, 10) optimizer = Adam(finetune_layer.parameters(), lr=1e-4) for epoch in range(10): # train_dataloader()加括号调用方法,才会返回可迭代的DataLoader对象 for batch in dm.train_dataloader(): x, y = batch with torch.no_grad(): features = backbone(x) features = features.flatten(1) preds = finetune_layer(features) loss = cross_entropy(preds, y) loss.backward() optimizer.step() optimizer.zero_grad() print(loss.item())
补充说明
- 如果需要加载验证集、测试集,调用
dm.val_dataloader()、dm.test_dataloader()即可,同样需要加括号调用方法 - 如果需要全流程使用数据集,
setup方法可以不传stage参数,直接调用dm.setup()即可完成所有阶段的数据集拆分
内容的提问来源于stack exchange,提问作者Wakame
相关产品推荐
相关产品推荐

