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

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)} "

原因分析

  1. 错误触发原因:inf和0都是违规值。PyTorch的Normal分布要求scale必须严格大于0,inf属于超出合理范围的无效值,0不满足GreaterThan(lower_bound=0.0)的严格大于要求,两者都会触发报错。
  2. 深层根源:
    • 从堆栈轨迹看,错误出现在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 03:54:53