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
相关产品推荐
相关产品推荐

