使用Stable Baselines3训练智能体时,Episode截断触发ValueError报错
自定义Gymnasium环境用Stable Baselines3训练时,截断(truncated)触发后抛出ValueError
报错信息
Traceback (most recent call last): File "C:\Users\bo112\PycharmProjects\ecocharge\code\Simulation Env\prototype_visu.py", line 684, in <module> model.learn(total_timesteps=time_steps, tb_log_name=log_name) File "C:\Users\bo112\PycharmProjects\ecocharge\venv\lib\site-packages\stable_baselines3\ppo\ppo.py", line 315, in learn return super().learn( File "C:\Users\bo112\PycharmProjects\ecocharge\venv\lib\site-packages\stable_baselines3\common\on_policy_algorithm.py", line 277, in learn continue_training = self.collect_rollouts(self.env, callback, self.rollout_buffer, n_rollout_steps=self.n_steps) File "C:\Users\bo112\PycharmProjects\ecocharge\venv\lib\site-packages\stable_baselines3\common\on_policy_algorithm.py", line 218, in collect_rollouts terminal_obs = self.policy.obs_to_tensor(infos[idx]["terminal_observation"])[0] File "C:\Users\bo112\PycharmProjects\ecocharge\venv\lib\site-packages\stable_baselines3\common\policies.py", line 256, in obs_to_tensor vectorized_env = vectorized_env or is_vectorized_observation(obs_, obs_space) File "C:\Users\bo112\PycharmProjects\ecocharge\venv\lib\site-packages\stable_baselines3\common\utils.py", line 399, in is_vectorized_observation return is_vec_obs_func(observation, observation_space) # type: ignore[operator] File "C:\Users\bo112\PycharmProjects\ecocharge\venv\lib\site-packages\stable_baselines3\common\utils.py", line 266, in is_vectorized_box_observation raise ValueError( ValueError: Error: Unexpected observation shape () for Box environment, please use (1,) or (n_env, 1) for the observation shape.
问题现象
在自定义Gymnasium环境中用Stable Baselines3的PPO训练智能体,程序仅在episode触发**截断(truncated)**时崩溃,触发终止(terminated)时正常。状态值生成逻辑未改动,不清楚观测形状为何变化,也不确定返回truncated和terminated是否有特殊规则。
相关代码
环境step函数
def step(self, action): ... # handling the action etc. reward = 0 truncated = False terminated = False # Check if time is over/score too low - else reward function if self.n_step >= self.max_steps: truncated = True print('truncated') elif self.score < -1000: terminated = True # print('terminated') else: reward = self.reward_fnc_distance() self.score += reward self.d_score.append(self.score) self.n_step += 1 # state: [current power, peak power, fridge 1 temp, fridge 2 temp, [...] , fridge n temp] self.state['current_power'] = self.d_power_sum[-1] self.state['peak_power'] = self.peak_power for i in range(self.n_fridges): self.state[f'fridge{i}_temp'] = self.d_fridges_temp[i][-1] self.state[f'fridge{i}_on'] = self.fridges[i].on if self.logging: print(f'score: {self.score}') if (truncated or terminated) and self.logging: self.save_run() return self.state, reward, terminated, truncated, {}
训练配置代码
hidden_layer = [64, 64, 32] time_steps = 1000_000 learning_rate = 0.003 log_name = f'PPO_{int(time_steps/1000)}k_lr{str(learning_rate).replace(".", "_")}' vec_env = make_vec_env(env_id=ChargeEnv, n_envs=4) model = PPO('MultiInputPolicy', vec_env, verbose=1, tensorboard_log='tensorboard_logs/', policy_kwargs={'net_arch': hidden_layer, 'activation_fn': th.nn.ReLU}, learning_rate=learning_rate, device=th.device("cuda" if th.cuda.is_available() else "cpu"), batch_size=128) model.learn(total_timesteps=time_steps, tb_log_name=log_name) model.save(f'models/{log_name}') vec_env.close()
解决方案
将self.state中所有float/Box类型的值转换为形状为(1,)的numpy数组后返回即可,修改后的状态赋值代码如下:
self.state['current_power'] = np.array([self.d_power_sum[-1]], dtype='float32') self.state['peak_power'] = np.array([self.peak_power], dtype='float32') for i in range(self.n_fridges): self.state[f'fridge{i}_temp'] = np.array([self.d_fridges_temp[i][-1]], dtype='float32') self.state[f'fridge{i}_on'] = self.fridges[i].on
注:指定dtype并非必须,但对于Stable Baselines3的SubprocVecEnv很重要。
内容的提问来源于stack exchange,提问作者maxxel_
相关产品推荐
相关产品推荐

