PyTorch Categorical性能优化求助:PPO智能体训练速度瓶颈
PyTorch下PPO中Categorical操作的提速方案
1. 跳过Categorical类实例化,手动实现核心逻辑
Categorical类的初始化会引入额外的封装与检查开销,直接用PyTorch原生张量操作实现采样、对数概率和熵计算,能大幅减少这部分耗时。
修改后的act函数:
def act(self, state): action_probs = self.actor(state) # 用multinomial直接采样,替代Categorical.sample() action = torch.multinomial(action_probs, num_samples=1, replacement=True) # 直接计算对数概率,跳过Categorical.log_prob() action_logprob = torch.log(action_probs.gather(1, action)) return action.detach(), action_logprob.detach()
修改后的evaluate函数:
def evaluate(self, state, action): action_probs = self.actor(state) # 计算对数概率 action_logprobs = torch.log(action_probs.gather(1, action)) # 手动计算熵:H = -Σ(p * log(p)),clamp避免log(0)的数值错误 dist_entropy = -(action_probs * torch.log(action_probs.clamp(min=1e-10))).sum(dim=-1) state_values = self.critic(state) return action_logprobs, state_values, dist_entropy
2. 改用Logits输出优化计算链
如果你的actor网络最后一层目前输出的是经过softmax的action_probs,可以改成输出未归一化的logits。这样能避免先softmax再取log的冗余计算,同时提升数值稳定性。
调整actor网络输出为logits后,修改后的函数:
def act(self, state): logits = self.actor(state) action_probs = torch.softmax(logits, dim=-1) action = torch.multinomial(action_probs, num_samples=1, replacement=True) # 用log_softmax直接得到对数概率,避免先softmax再log action_logprob = torch.nn.functional.log_softmax(logits, dim=-1).gather(1, action) return action.detach(), action_logprob.detach() def evaluate(self, state, action): logits = self.actor(state) log_probs = torch.nn.functional.log_softmax(logits, dim=-1) action_logprobs = log_probs.gather(1, action) # 基于logits计算熵,复用已计算的log_probs减少冗余 action_probs = torch.softmax(logits, dim=-1) dist_entropy = -(action_probs * log_probs).sum(dim=-1) state_values = self.critic(state) return action_logprobs, state_values, dist_entropy
注:可以把action_probs缓存下来复用,进一步减少重复计算量。
3. 确保批量操作最大化效率
确认所有输入输出都保持批量维度(比如state是[batch_size, state_dim],action是[batch_size, 1]),PyTorch的批量张量操作比逐样本处理快一个数量级以上。避免在循环中处理单个样本,所有逻辑都基于批量张量实现。
内容的提问来源于stack exchange,提问作者Erik Storm
相关产品推荐
相关产品推荐

