使用Ray包时执行loss.backward未更新PyTorch模型权重问题求助
问题根因
你遇到的权重不更新、损失无下降问题的核心原因是PyTorch自动微分计算图在跨Ray进程传递时断裂:
- Ray的远程任务(
@ray.remote装饰的函数)接收张量参数时,会将主进程的张量序列化后传输到worker进程,该过程仅保留张量的数值,丢失了和主进程模型关联的grad_fn梯度链路 - 最终通过
ray.get拿到的损失值是无梯度属性的普通标量,调用loss.backward()时没有梯度可以回传到模型参数,因此权重不会更新。
此外你提供的测试代码还存在两个可优化点:
- 缺少
ray.init()初始化语句,无法正常启动Ray runtime - DNN类中重复定义了两次
self.relu,属于冗余代码
修复方案
我们可以通过自定义PyTorch Autograd Function封装Ray的远程调用,手动维护前向和反向的梯度传递链路,保证计算图的完整性,同时保留你原本的分数GPU并行能力。
修复后的完整可运行代码如下:
import torch import ray # 初始化Ray runtime,按你的GPU实际数量调整num_gpus参数 ray.init(num_gpus=1) training_set = torch.FloatTensor([[1,1],[2,2],[3,3]]) class DNN(torch.nn.Module): def __init__(self): super(DNN,self).__init__() self.linear_1=torch.nn.Linear(2,3) self.relu1=torch.nn.ReLU() self.linear_2=torch.nn.Linear(3,3) self.relu2=torch.nn.ReLU() self.linear_3=torch.nn.Linear(3,2) def forward(self,input_tensor): linear1=self.linear_1(input_tensor) relu1=self.relu1(linear1) linear2=self.linear_2(relu1) relu2=self.relu2(linear2) linear3=self.linear_3(relu2) output=linear3 return output # 自定义Autograd Function封装Ray远程损失计算 class ParallelLossFunction(torch.autograd.Function): @staticmethod def forward(ctx, x_ray): ctx.save_for_backward(x_ray) x_len = len(x_ray) parallel_out = ray.get([custom_loss.remote(x, idx) for idx, x in enumerate(x_ray)]) return sum(parallel_out)/x_len @staticmethod def backward(ctx, grad_output): x_ray, = ctx.saved_tensors # 匹配你的x平方损失的反向梯度计算,可根据实际损失逻辑修改 grad_input = 2 * x_ray * grad_output / len(x_ray) return grad_input @ray.remote(num_gpus=0.25) def custom_loss(x, idx): print(x,idx) return (x[0]**2 + x[1]**2) def parallelize_loss(x_ray): return ParallelLossFunction.apply(x_ray) model=DNN() learning_rate=0.1 epochs=10 # 调大epoch方便观察损失下降 optimizer=torch.optim.Adam(model.parameters(),lr=learning_rate) for epoch in range(epochs): model.train() optimizer.zero_grad() train_output=model(training_set) loss = parallelize_loss(train_output) print("loss at",epoch,"=",loss.item()) loss.backward() optimizer.step() # 关闭Ray ray.shutdown()
验证效果
运行修复后的代码可以看到损失值逐轮下降,模型权重正常更新,同时会启动4个占用0.25GPU的worker执行损失计算任务,符合你的需求。
内容的提问来源于stack exchange,提问作者zuki
相关产品推荐
相关产品推荐

