使用torch.nn.DataParallel多GPU训练时模型参数无法更新如何解决?
问题:多GPU下模型成员变量无法更新,单GPU正常
复现代码
import torch import torch.nn as nn import os class Net(nn.Module): def __init__(self): super().__init__() self.h = -1 def forward(self, x): self.h =x os.environ['CUDA_VISIBLE_DEVICES'] = '0' if torch.cuda.is_available(): print('using Cuda devices, num:', torch.cuda.device_count()) model = nn.DataParallel(Net()) x = 2 print(model.module.h) model(x) print(model.module.h)
现象
- 双GPU运行:
h始终保持初始值-1,未被更新 - 单GPU运行:
h成功从-1更新为2
原因
nn.DataParallel会为每个GPU创建独立的模型副本,调用model(x)时,输入会被拆分到各个GPU副本执行forward,但只有主GPU的计算结果会返回,而你直接访问的model.module是原始的模型实例,并非执行forward的副本。副本中的h确实被更新了,但这些变更不会自动同步回原始实例,导致你看到的model.module.h始终是初始值。
解决方案
方法1:用register_buffer存储状态(推荐)
将需要跨GPU同步的状态注册为模型的缓冲区,这样DataParallel会自动同步各个副本的状态到主模型:
import torch import torch.nn as nn import os class Net(nn.Module): def __init__(self): super().__init__() # 注册缓冲区,不需要梯度计算 self.register_buffer('h', torch.tensor(-1, dtype=torch.int)) def forward(self, x): # 用copy_方法更新缓冲区,保证数据同步 self.h.copy_(torch.tensor(x, dtype=torch.int, device=self.h.device)) os.environ['CUDA_VISIBLE_DEVICES'] = '0,1' # 双GPU示例 if torch.cuda.is_available(): print('using Cuda devices, num:', torch.cuda.device_count()) model = nn.DataParallel(Net()) x = 2 print(model.module.h.item()) model(x) print(model.module.h.item())
方法2:改用nn.Parameter(如果需要梯度)
如果这个状态需要参与梯度计算,改用nn.Parameter:
class Net(nn.Module): def __init__(self): super().__init__() self.h = nn.Parameter(torch.tensor(-1, dtype=torch.int)) def forward(self, x): self.h.data.copy_(torch.tensor(x, dtype=torch.int, device=self.h.device))
说明
用register_buffer或nn.Parameter的核心是让变量纳入模型的状态管理体系,DataParallel会自动处理跨GPU的状态同步,确保你访问model.module时能拿到更新后的值。
内容的提问来源于stack exchange,提问作者hescluke
相关产品推荐
相关产品推荐

