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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.07 14:51:03