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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.23 05:33:17