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

PyTorch复现Actor-Critic代码遇RuntimeError问题求助

修复RuntimeError:重复反向传播计算图的问题

问题根源

这个报错的核心是同一个计算图被多次执行反向传播,且中间张量的计算图未被正确分离。你的代码里,actor和critic的更新流程共享了带梯度的张量,第一次反向传播(actor更新)后,PyTorch会自动释放计算图资源,第二次(critic更新)再尝试复用这些资源就会触发错误。

具体修复步骤

1. 分离计算图,切断梯度传播链路

计算TD目标和TD误差时,必须用.detach()把critic的输出从计算图中分离出来,避免后续反向传播互相干扰:

# 训练循环中替换TD目标和误差的计算逻辑
next_state_value = critic(next_state).detach().squeeze().item()
target = reward + gamma * next_state_value

current_state_value_detached = critic(state).detach().squeeze().item()
td_error = target - current_state_value_detached
td_error = torch.tensor(td_error, dtype=torch.float32)

2. 修正Critic的更新逻辑,避免重复使用计算图张量

原代码中直接把critic(state)的带梯度结果传入compute_return,会导致反向传播时复用已释放的图资源。修改为:

  • 训练循环中单独获取带梯度的当前状态值,再传入更新方法
  • 确保target转为张量类型,匹配模型输出的格式

训练循环中更新critic的代码改为:

current_state_value_with_grad = critic(state).squeeze()
critic.compute_return(current_state_value_with_grad, target)

Critic类的compute_return方法修改为:

def compute_return(self, output, target):
    self.optimizer.zero_grad()
    target = torch.tensor(target, dtype=torch.float32)
    loss = self.criterion(output, target)
    loss.backward()
    self.optimizer.step()

3. 修复Policy网络的标准差计算错误

原代码用F.softmax处理标准差sigma是逻辑错误——softmax会让输出和为1,但标准差必须是正数,应该用F.softplus确保sigma为正:

def forward(self, state):
    state = torch.from_numpy(state).float()
    x = F.relu(self.hidden1(state))
    x = F.relu(self.hidden2(x))
    mu = self.mu(x)
    # 替换softmax为softplus,保证标准差为正数
    sigma = F.softplus(self.sigma(x))
    action_dist = Normal(mu, sigma)
    action_var = action_dist.rsample()
    action_var = torch.clip(action_var,
                    self.env.action_space.low[0],
                    self.env.action_space.high[0])

    return action_var, action_dist

4. 启用状态归一化(关键收敛优化)

你已经定义了scale_state函数,但训练循环中完全没用到,状态未归一化会严重影响模型收敛,记得在获取状态后调用:

# 训练循环中初始化状态时
state = scale_state(env.reset())
# 每次获取next_state后
next_state = scale_state(next_state)

额外说明

  • 移除不必要的重复zero_grad():optimizer.step()后不需要立刻调用zero_grad(),下次更新前调用即可,重复操作不会解决报错问题,反而多余。

内容的提问来源于stack exchange,提问作者EmptySet

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 04:36:22