Stable Baselines3训练MLP A2C策略时遇Normal分布参数约束报错
问题与解决方案
问题详情
核心报错
ValueError('Expected parameter scale (Tensor of shape (1, 4)) of distribution Normal(loc: torch.Size([1, 4]), scale: torch.Size([1, 4])) to satisfy the constraint GreaterThan(lower_bound=0.0), but found invalid values:
tensor([[inf, inf, 0., 0.]])')
背景信息
- 动作维度:
(4,),观测维度:(3,) - 使用Stable Baselines3的
model.learn训练,前期正常,权重更新阶段触发错误 - 疑问:
- 错误是
inf还是0导致? - 是否需要调整动作范围为
0<a<=1? - 是否是Stable Baselines3包的问题?
- 错误是
报错堆栈
~\anaconda3\envs\\lib\site-packages\stable_baselines3\common\on_policy_algorithm.py in learn(self, total_timesteps, callback, log_interval, tb_log_name, reset_num_timesteps, progress_bar) 257 258 while self.num_timesteps < total_timesteps: --> 259 continue_training = self.collect_rollouts(self.env, callback, self.rollout_buffer, n_rollout_steps=self.n_steps) 260 261 if continue_training is False: ~\anaconda3\envs\\lib\site-packages\stable_baselines3\common\on_policy_algorithm.py in collect_rollouts(self, env, callback, rollout_buffer, n_rollout_steps) 167 # Convert to pytorch tensor or to TensorDict 168 obs_tensor = obs_as_tensor(self._last_obs, self.device) --> 169 actions, values, log_probs = self.policy(obs_tensor) 170 actions = actions.cpu().numpy() 171 ~\anaconda3\envs\\lib\site-packages\torch\nn\modules\module.py in _call_impl(self, *input, **kwargs) 1192 if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks 1193 or _global_forward_hooks or _global_forward_pre_hooks): -> 1194 return forward_call(*input, **kwargs) 1195 # Do not call functions when jit is used 1196 full_backward_hooks, non_full_backward_hooks = [], [] ~\anaconda3\envs\\lib\site-packages\stable_baselines3\common\policies.py in forward(self, obs, deterministic) 624 # Evaluate the values for the given observations 625 values = self.value_net(latent_vf) --> 626 distribution = self._get_action_dist_from_latent(latent_pi) 627 actions = distribution.get_actions(deterministic=deterministic) 628 log_prob = distribution.log_prob(actions) ~\anaconda3\envs\\lib\site-packages\stable_baselines3\common\policies.py in _get_action_dist_from_latent(self, latent_pi) 654 655 if isinstance(self.action_dist, DiagGaussianDistribution): --> 656 return self.action_dist.proba_distribution(mean_actions, self.log_std) 657 elif isinstance(self.action_dist, CategoricalDistribution): 658 # Here mean_actions are the logits before the softmax ~\anaconda3\envs\\lib\site-packages\stable_baselines3\common\distributions.py in proba_distribution(self, mean_actions, log_std) 162 """ 163 action_std = th.ones_like(mean_actions) * log_std.exp() --> 164 self.distribution = Normal(mean_actions, action_std) 165 return self 166 ~\anaconda3\envs\\lib\site-packages\torch\distributions\normal.py in __init__(self, loc, scale, validate_args) 54 else: 55 batch_shape = self.loc.size() --> 56 super(Normal, self).__init__(batch_shape, validate_args=validate_args) 57 58 def expand(self, batch_shape, _instance=None): ~\anaconda3\envs\\lib\site-packages\torch\distributions\distribution.py in __init__(self, batch_shape, event_shape, validate_args) 55 if not valid.all(): 56 raise ValueError( --> 57 f"Expected parameter {param} " 58 f"({type(value).__name__} of shape {tuple(value.shape)}) " 59 f"of distribution {repr(self)} "
原因分析
- 错误触发原因:
inf和0都是违规值。PyTorch的Normal分布要求scale必须严格大于0,inf属于超出合理范围的无效值,0不满足GreaterThan(lower_bound=0.0)的严格大于要求,两者都会触发报错。 - 深层根源:
- 从堆栈轨迹看,错误出现在
DiagGaussianDistribution生成正态分布时,log_std.exp()得到了异常值。log_std是模型的可学习参数,训练过程中梯度爆炸或消失会导致其值突变:- 若
log_std过大(如远大于20),exp()后会溢出为inf; - 若
log_std过小(如远小于-20),exp()后会趋近于0。
- 若
- 该问题和动作范围
0<=a<=1无关,也不是Stable Baselines3包的固有问题,属于训练过程中的参数更新异常。
- 从堆栈轨迹看,错误出现在
解决办法
- 限制
log_std取值范围:初始化模型时,通过policy_kwargs设置log_std的上下界,避免exp()后出现inf或0。示例:
这里model = PPO("MlpPolicy", env, policy_kwargs={"log_std_init": 0.0, "log_std_range": (-20, 20)}, verbose=1)log_std_range将参数限制在-20到20之间,exp(-20)约为2e-9(严格大于0),exp(20)约为4.8e8(不会溢出为inf)。 - 降低学习率:若梯度爆炸导致
log_std突变,尝试降低学习率,比如将默认的3e-4调整为1e-4或更小。 - 启用梯度裁剪:通过
policy_kwargs设置梯度裁剪,防止参数更新幅度过大。示例:policy_kwargs = dict(optimizer_kwargs={"max_grad_norm": 0.5}) model = PPO("MlpPolicy", env, policy_kwargs=policy_kwargs, verbose=1) - 检查奖励信号:如果环境奖励波动过大(如突然出现极大/极小值),会引发梯度异常,需平滑奖励或优化奖励函数设计。
内容的提问来源于stack exchange,提问作者Snared
相关产品推荐
相关产品推荐

