From a60cc188740debd45def5a3512f867164f64d333 Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Sat, 22 Jan 2022 08:05:16 +0100 Subject: [PATCH] PPO lidar + random agent --- .../intersimple/ppo_speed_lidar_random.py | 49 +++++++++++++++++++ 1 file changed, 49 insertions(+) create mode 100644 scratch/etienne/intersimple/ppo_speed_lidar_random.py diff --git a/scratch/etienne/intersimple/ppo_speed_lidar_random.py b/scratch/etienne/intersimple/ppo_speed_lidar_random.py new file mode 100644 index 0000000..0da4cc5 --- /dev/null +++ b/scratch/etienne/intersimple/ppo_speed_lidar_random.py @@ -0,0 +1,49 @@ +# %% +from stable_baselines3 import PPO +from intersim.envs import IntersimpleLidarFlatRandom +from intersim.envs.intersimple import speed_reward +import functools + +model_name = "ppo_speed_lidar_random" + +#def reward(state, action, info): +# speed = state[2].item() +# r = speed if speed < 10 else (10 - 5 * (speed - 10)) +# return 0.1 * r + +env = IntersimpleLidarFlatRandom( + n_rays=5, + reward=functools.partial( + speed_reward, + collision_penalty=0 + ), +) + +# %% +model = PPO( + "MlpPolicy", env, + learning_rate=1e-4, + verbose=1, + tensorboard_log='runs/' +) +model.learn(total_timesteps=1000000) +model.save(model_name) + +print('Done training.') + +del model # remove to demonstrate saving and loading + +# %% +model = PPO.load(model_name) + +obs = env.reset() +while True: + action, _states = model.predict(obs) + obs, rewards, done, info = env.step(action) + env.render(mode='post') + if done: + break + +env.close(filestr='render/'+model_name) + +# %%