11from __future__ import annotations
22
3+ from pathlib import Path
34from typing import Any
45
56import gymnasium as gym
67import matplotlib .pyplot as plt
78from stable_baselines3 import PPO
89from stable_baselines3 .common .env_util import make_vec_env
910from stable_baselines3 .common .evaluation import evaluate_policy
11+ from stable_baselines3 .common .vec_env import VecNormalize
1012
1113plt .rcParams ["figure.figsize" ] = (10 , 5 )
1214
@@ -17,6 +19,7 @@ class ProximalPolicyOptimizationAgent:
1719 Initializes an agent that learns a policy via PPO algorithm to solve the task at hand.
1820
1921 PPO agents can be saved as zip file and re-loaded to avoid re-training.
22+ VecNormalize statistics are saved alongside the model as `<name>_vecnorm.pkl`.
2023
2124 Args:
2225 env (gym.Env): the environment the agent is acting on.
@@ -34,43 +37,62 @@ def __init__(
3437 trained : tuple [str , bool ] | None = None ,
3538 ):
3639 self .trained = trained
37- if env_kwargs is None :
38- self .env = env () # type: ignore[operator] ## the object is callable! (__init__())
39- else :
40- self .env = env (** env_kwargs ) # type: ignore[operator] ## the object is callable! (__init__())
41- _n_envs = n_envs = 1 if n_envs <= 0 else n_envs
42- self .vec_env = make_vec_env (env_id = env , n_envs = _n_envs , env_kwargs = env_kwargs ) # type: ignore ## should be correct
43- if n_envs <= 0 :
44- assert self .trained is not None , "When no training is specified a saved model should be provided"
45- self .model = PPO .load (self .trained [0 ])
46- elif n_envs == 1 :
47- self .model = PPO ("MlpPolicy" , self .env , verbose = 1 )
48- if trained is not None :
49- self .trained = (trained [0 ], trained [1 ])
40+ inference_only = n_envs <= 0
41+ _n_envs = 1 if inference_only else n_envs
42+
43+ raw_vec_env = make_vec_env (env_id = env , n_envs = _n_envs , env_kwargs = env_kwargs ) # type: ignore
44+
45+ if inference_only :
46+ assert trained is not None , "When no training is specified a saved model should be provided"
47+ stats_path = self ._stats_path (trained [0 ])
48+ if stats_path .exists ():
49+ self .vec_env = VecNormalize .load (str (stats_path ), raw_vec_env )
50+ else :
51+ self .vec_env = VecNormalize (raw_vec_env , norm_obs = True , norm_reward = False )
52+ self .vec_env .training = False
53+ self .vec_env .norm_reward = False
54+ self .model = PPO .load (trained [0 ], env = self .vec_env )
5055 else :
51- self .model = PPO ("MlpPolicy" , self .vec_env )
56+ self .vec_env = VecNormalize (raw_vec_env , norm_obs = True , norm_reward = True )
57+ if _n_envs == 1 :
58+ self .model = PPO ("MlpPolicy" , self .vec_env , verbose = 1 )
59+ else :
60+ self .model = PPO ("MlpPolicy" , self .vec_env )
5261 self .trained = (
53- trained [0 ] if trained is not None else f"ppo_{ env .__name__ } " , # type: ignore ## has name
62+ trained [0 ] if trained is not None else f"ppo_{ env .__name__ } " , # type: ignore[attr-defined]
5463 False if trained is None else trained [1 ],
5564 )
5665
66+ # Single unwrapped env for do_one_episode/evaluate without reconstructing a new crane.
67+ self .env = self .vec_env .venv .envs [0 ] # type: ignore[attr-defined]
68+
69+ @staticmethod
70+ def _stats_path (model_path : str ) -> Path :
71+ p = Path (model_path )
72+ return p .parent / f"{ p .stem } _vecnorm.pkl"
73+
5774 def do_training (self , total_timesteps : int = 25000 , progress_bar : bool = True ):
5875 self .model .learn (total_timesteps , progress_bar = progress_bar )
59- if self .trained is not None and self .trained [1 ] and self .env .render_mode not in ("play-back" ):
76+ if self .trained is not None and self .trained [1 ] and self .env .render_mode not in ("play-back" , ):
6077 self .model .save (self .trained [0 ])
78+ self .vec_env .save (str (self ._stats_path (self .trained [0 ])))
6179
6280 def evaluate (self , n_episodes : int = 10 ):
63- mean_reward , std_reward = evaluate_policy (self .model , self .env , n_eval_episodes = n_episodes )
81+ self .vec_env .training = False
82+ self .vec_env .norm_reward = False
83+ mean_reward , std_reward = evaluate_policy (self .model , self .vec_env , n_eval_episodes = n_episodes )
84+ self .vec_env .training = True
85+ self .vec_env .norm_reward = True
6486 print (f"Mean:{ mean_reward } , stdev:{ std_reward } " )
6587
6688 def do_one_episode (self , seed : int = 1 ):
67- """Do one episode on the non-vectorized, trained environment ."""
89+ """Do one episode using the trained normalizer for observations ."""
6890 obs , info = self .env .reset (seed = seed )
6991 terminated = truncated = False
7092 while not terminated and not truncated :
71- action , _states = self .model .predict (obs )
93+ norm_obs = self .vec_env .normalize_obs (obs )
94+ action , _states = self .model .predict (norm_obs , deterministic = True )
7295 obs , rewards , terminated , truncated , info = self .env .step (action )
73- # print("Action", obs, rewards, terminated, truncated, info)
7496 self .env .render ()
7597
7698
0 commit comments