多离散动作空间下简易Actor-Critic策略网络更新正确性验证求助
多离散动作空间Actor-Critic策略更新验证结论
你的适配多离散动作空间的Actor-Critic策略网络更新实现是正确的,核心逻辑完全符合策略梯度算法在多离散动作场景下的要求,以下是具体验证和优化建议:
核心逻辑正确性说明
多离散动作空间中,每个动作分支是独立的离散分布,智能体选择的动作联合概率是各分支动作概率的乘积。根据策略梯度的对数技巧,目标函数的对数形式就是各分支动作对数概率的和。你的代码中:
sum(selected_log_probs)正确计算了联合动作的对数概率之和- 乘以
td_delta.detach()(优势函数,且阻断梯度流向Critic),再取负均值作为损失,等价于最大化期望回报的策略梯度目标,逻辑完全正确。
代码细节验证
概率分布与对数概率获取:
output_probs = self.actor(states) log_probs = [torch.log(prob) for prob in output_probs]正确获取了每个动作分支的softmax概率分布及其对数形式,符合多离散动作的输出设计。
选中动作的对数概率提取:
selected_log_probs.append(log_probs[i].gather(1, actions[:, i].unsqueeze(1)))通过
gather操作准确取出每个样本对应动作的对数概率,unsqueeze(1)保证了张量维度匹配,避免广播错误。损失计算与梯度隔离:
actor_loss = torch.mean(-sum(selected_log_probs) * td_delta.detach())使用
td_delta.detach()确保Critic的梯度不会影响Actor的更新,符合Actor-Critic双网络分离训练的原则。
可优化的细节点
1. 批量张量操作替代循环
原循环可以替换为更高效的批量操作,提升训练速度:
# 替代原循环中的log_probs和selected_log_probs构建 log_probs = torch.cat([torch.log(prob) for prob in output_probs], dim=1) selected_log_probs = log_probs.gather(1, actions) actor_loss = torch.mean(-selected_log_probs.sum(dim=1, keepdim=True) * td_delta.detach())
2. 设备一致性修复
PolicyNet.forward中torch.cumsum创建的张量默认在CPU,若模型部署在GPU会导致维度不匹配,修改为:
def forward(self, x): x = self.fc1(x) x = F.relu(x) x = self.fc2(x) x = F.relu(x) x = self.fc3(x) # 确保维度起始张量与输入同设备 dim_starts = torch.cumsum(torch.tensor([0] + self.action_dims[:-1], device=x.device), dim=0) out = [F.softmax(x[:, 0:self.action_dims[0]], dim=-1)] \ + [F.softmax(x[:, s:s+d], dim=-1) for s, d in zip(dim_starts[1:], self.action_dims[1:])] return out
3. 状态张量创建优化
take_action中的状态转换可以简化并适配设备:
def take_action(self, state): # 自动匹配模型设备,避免numpy转换 state = torch.as_tensor(state, dtype=torch.float32, device=next(self.actor.parameters()).device).unsqueeze(0) action_probs = self.actor(state) actions = [torch.multinomial(prob, 1, replacement=True).item() for prob in action_probs] return actions
内容的提问来源于stack exchange,提问作者Reese
相关产品推荐
相关产品推荐

