PyTorch视频分类训练报错AttributeError: 无法在__main__模块找到video_dataset属性
错误产生原因
该报错为Windows系统下Python multiprocessing 模块采用spawn启动方式的典型问题:
- PyTorch DataLoader设置
num_workers>0时会启动子进程加载数据,子进程启动时需序列化(pickle)导入主模块的所有内容,包括自定义的video_dataset类 - 若
if __name__ == '__main__'包裹范围错误,或是video_dataset类定义放在了if分支内部,会导致子进程加载主模块时无法找到该类的定义 - 若使用
import_ipynb导入ipynb文件的代码,跨文件导入自定义类时也容易出现序列化找不到类的问题 video_dataset类的__getitem__方法中调用了全局变量im_size、decoder,序列化时也可能出现全局变量无法找到的问题
解决方案
- 调整
if __name__ == '__main__'的包裹范围:所有导入语句、类定义、函数定义、全局常量定义全部放在if分支外部,仅将实例化数据集、实例化DataLoader、执行训练的逻辑放在if分支内部,参考代码结构如下:
# 所有导入放在最外层 import torch from torch.utils.data import Dataset, DataLoader import cv2 import numpy as np # 其余依赖导入省略 # 自定义类、函数定义放在最外层 class video_dataset(Dataset): def __init__(self,frame_list,sequence_length = 16,transform = None, im_size=224, decoder=None): self.frame_list = frame_list self.transform = transform self.sequence_length = sequence_length # 将全局变量改为类属性传入,避免序列化异常 self.im_size = im_size self.decoder = decoder def __len__(self): return len(self.frame_list) def __getitem__(self,idx): label,path = self.frame_list[idx] img = cv2.imread(path) seq_img = list() for i in range(16): img1 = img[:,128*i:128*(i+1),:] if(self.transform): img1 = self.transform(img1) seq_img.append(img1) seq_image = torch.stack(seq_img) seq_image = seq_image.reshape(3,16,self.im_size,self.im_size) return seq_image, self.decoder[label] # 其他全局常量、transform、函数定义放在外层 im_size = 224 decoder = {"class1":0, "class2":1} # 替换为实际的标签映射 # 其余全局定义省略 # 仅执行逻辑放在if分支内 if __name__ == '__main__': # 实例化数据集时传入之前的全局变量 train_dataset = video_dataset(train_frame_list, im_size=im_size, decoder=decoder, transform=train_transform) val_dataset = video_dataset(val_frame_list, im_size=im_size, decoder=decoder, transform=val_transform) # 实例化DataLoader train_loader = DataLoader(train_dataset, batch_size=8, num_workers=4, shuffle=True) val_loader = DataLoader(val_dataset, batch_size=8, num_workers=4, shuffle=False) dataloaders = {"train": train_loader, "val": val_loader} # 训练逻辑 from torch.autograd import Variable iteration = 0 acc_all = list() loss_all = list() num_epochs = 20 # 替换为实际设置的训练轮数 # 其余训练相关初始化(model、optimizer、损失函数等)省略 for epoch in range(num_epochs): # 原有训练循环代码 ...
- 调整后如果仍然报错,可临时将DataLoader的
num_workers参数设置为0,禁用多进程数据加载,验证问题来源后再进一步排查类作用域问题 - 若在Jupyter Notebook中运行代码,Windows下多进程兼容性较差,可将代码导出为
.py文件后通过命令行执行,不要直接在notebook中运行带多进程数据加载的训练逻辑
内容的提问来源于stack exchange,提问作者MinJae
相关产品推荐
相关产品推荐

