使用PyTorch DataParallel多GPU训练时无法调用原自定义模型函数如何解决?
解决nn.DataParallel封装后无法调用原自定义模型函数的方案
- 方法1:直接访问
module属性nn.DataParallel会将原始模型存储在实例的.module属性下,调用自定义方法时直接访问该属性即可,示例如下:
# 假设原自定义模型有custom_process()方法 # 封装前调用方式 model.custom_process() # 封装后调用方式 model.module.custom_process()
该方法无需修改原有代码结构,适合临时调用自定义方法的场景。
- 方法2:保存/加载模型时直接操作原始模型,避免封装结构干扰
训练时保存模型权重,仅保存原始模型的state_dict,而非整个DataParallel实例,后续加载权重时直接加载到原始模型实例上,无需处理DataParallel封装层,可正常调用自定义方法:
# 训练时保存权重 torch.save(model.module.state_dict(), "model_weights.pth") # 后续加载使用 model = YourCustomModel() model.load_state_dict(torch.load("model_weights.pth")) model.custom_process() # 可直接调用
- 方法3:自定义封装类转发常用方法(可选)
如果需要频繁调用自定义方法,不想每次都编写.module前缀,可以自定义封装类自动转发指定方法:
class CustomDataParallel(nn.DataParallel): def custom_process(self, *args, **kwargs): return self.module.custom_process(*args, **kwargs) # 封装时使用自定义类 model = CustomDataParallel(model) model.custom_process() # 可直接调用
目前PyTorch官方已不推荐使用
nn.DataParallel,其单进程多线程的模式存在GIL锁性能瓶颈,更推荐使用nn.parallel.DistributedDataParallel实现多卡训练。DDP同样通过.module属性存储原始模型,上述所有解决方案对DDP同样适用。
内容的提问来源于stack exchange,提问作者pandas
相关产品推荐
相关产品推荐

