Port TRPO, PPO, GAIL
This commit is contained in:
55
scratch/etienne/trpo/sb3/sb3-ppo-intersimple-nocollision.py
Normal file
55
scratch/etienne/trpo/sb3/sb3-ppo-intersimple-nocollision.py
Normal file
@@ -0,0 +1,55 @@
|
||||
from stable_baselines3 import PPO
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from gym import Wrapper
|
||||
|
||||
model_name = "ppo_speed_lidar_nocollision"
|
||||
|
||||
env = IntersimpleLidarFlat(
|
||||
n_rays=5,
|
||||
agent=51,
|
||||
reward=functools.partial(
|
||||
speed_reward,
|
||||
collision_penalty=0
|
||||
),
|
||||
stop_on_collision=False,
|
||||
)
|
||||
|
||||
class CollisionPenaltyWrapper(Wrapper):
|
||||
|
||||
def __init__(self, env, collision_distance, collision_penalty, *args, **kwargs):
|
||||
super().__init__(env, *args, **kwargs)
|
||||
self.penalty = collision_penalty
|
||||
self.distance = collision_distance
|
||||
|
||||
def step(self, action):
|
||||
obs, reward, done, info = super().step(action)
|
||||
reward = -self.penalty if (obs.reshape(-1, 6)[1:, 0] < self.distance).any() else reward
|
||||
|
||||
self.env._rewards.pop()
|
||||
self.env._rewards.append(reward)
|
||||
|
||||
return obs, reward, done, info
|
||||
|
||||
env = CollisionPenaltyWrapper(env, collision_distance=6, collision_penalty=100)
|
||||
|
||||
model = PPO(
|
||||
"MlpPolicy", env,
|
||||
learning_rate=1e-4,
|
||||
verbose=1,
|
||||
)
|
||||
model.learn(total_timesteps=100000)
|
||||
model.save(model_name)
|
||||
|
||||
model = PPO.load(model_name)
|
||||
obs = env.reset()
|
||||
env.render(mode='post')
|
||||
for i in range(200):
|
||||
action, _ = model.predict(obs)
|
||||
obs, reward, done, _ = env.step(action)
|
||||
env.render(mode='post')
|
||||
print('step', i, 'front distance', obs.reshape(-1, 6)[3, 0], 'reward', reward)
|
||||
if done:
|
||||
break
|
||||
env.close()
|
||||
58
scratch/etienne/trpo/sb3/sb3-ppo-intersimple-nocollision2.py
Normal file
58
scratch/etienne/trpo/sb3/sb3-ppo-intersimple-nocollision2.py
Normal file
@@ -0,0 +1,58 @@
|
||||
from stable_baselines3 import PPO
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from gym import Wrapper
|
||||
|
||||
model_name = "ppo_speed_lidar_nocollision"
|
||||
|
||||
env = IntersimpleLidarFlat(
|
||||
n_rays=5,
|
||||
agent=51,
|
||||
reward=functools.partial(
|
||||
speed_reward,
|
||||
collision_penalty=0
|
||||
),
|
||||
stop_on_collision=False,
|
||||
)
|
||||
|
||||
class CollisionPenaltyWrapper(Wrapper):
|
||||
|
||||
def __init__(self, env, collision_distance, collision_penalty, last_reward_weight, *args, **kwargs):
|
||||
super().__init__(env, *args, **kwargs)
|
||||
self.penalty = collision_penalty
|
||||
self.distance = collision_distance
|
||||
self.last_reward = -collision_penalty
|
||||
self.last_reward_weight = last_reward_weight
|
||||
|
||||
def step(self, action):
|
||||
obs, reward, done, info = super().step(action)
|
||||
reward = -self.penalty if (obs.reshape(-1, 6)[1:, 0] < self.distance).any() else reward
|
||||
reward = self.last_reward_weight * self.last_reward + (1 - self.last_reward_weight) * self.last_reward
|
||||
|
||||
self.env._rewards.pop()
|
||||
self.env._rewards.append(reward)
|
||||
|
||||
return obs, reward, done, info
|
||||
|
||||
env = CollisionPenaltyWrapper(env, collision_distance=6, collision_penalty=10, last_reward_weight=0.9)
|
||||
|
||||
model = PPO(
|
||||
"MlpPolicy", env,
|
||||
learning_rate=1e-4,
|
||||
verbose=1,
|
||||
)
|
||||
model.learn(total_timesteps=100000)
|
||||
model.save(model_name)
|
||||
|
||||
model = PPO.load(model_name)
|
||||
obs = env.reset()
|
||||
env.render(mode='post')
|
||||
for i in range(200):
|
||||
action, _ = model.predict(obs)
|
||||
obs, reward, done, _ = env.step(action)
|
||||
env.render(mode='post')
|
||||
print('step', i, 'front distance', obs.reshape(-1, 6)[3, 0], 'reward', reward)
|
||||
if done:
|
||||
break
|
||||
env.close()
|
||||
28
scratch/etienne/trpo/sb3/sb3-ppo-intersimple-rollout.py
Normal file
28
scratch/etienne/trpo/sb3/sb3-ppo-intersimple-rollout.py
Normal file
@@ -0,0 +1,28 @@
|
||||
import sys
|
||||
sys.path.append('..')
|
||||
|
||||
from stable_baselines3 import PPO
|
||||
from core.sampling import rollout_sb3
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
import torch
|
||||
from wrappers import CollisionPenaltyWrapper
|
||||
|
||||
model = PPO.load('sb3-ppo-intersimple')
|
||||
env = CollisionPenaltyWrapper(IntersimpleLidarFlat(
|
||||
n_rays=5,
|
||||
agent=51,
|
||||
reward=functools.partial(
|
||||
speed_reward,
|
||||
collision_penalty=0
|
||||
),
|
||||
), collision_distance=6, collision_penalty=100)
|
||||
|
||||
expert_data = rollout_sb3(env, model, n_episodes=200, max_steps_per_episode=200)
|
||||
|
||||
states, actions, rewards, dones = expert_data
|
||||
print(f'Expert mean episode length {(~dones).sum() / states.shape[0]}')
|
||||
print(f'Expert mean reward per episode {rewards[~dones].sum() / states.shape[0]}')
|
||||
|
||||
torch.save(expert_data, 'sb3-ppo-intersimple-expert-data.pt')
|
||||
23
scratch/etienne/trpo/sb3/sb3-ppo-intersimple.py
Normal file
23
scratch/etienne/trpo/sb3/sb3-ppo-intersimple.py
Normal file
@@ -0,0 +1,23 @@
|
||||
from stable_baselines3 import PPO
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
|
||||
env = IntersimpleLidarFlat(
|
||||
n_rays=5,
|
||||
agent=51,
|
||||
reward=functools.partial(
|
||||
speed_reward,
|
||||
collision_penalty=0
|
||||
),
|
||||
)
|
||||
|
||||
model = PPO(
|
||||
"MlpPolicy", env,
|
||||
learning_rate=1e-4,
|
||||
verbose=1,
|
||||
use_sde=False,
|
||||
sde_sample_freq=4,
|
||||
)
|
||||
model.learn(total_timesteps=100000)
|
||||
model.save('sb3-ppo-intersimple')
|
||||
6
scratch/etienne/trpo/sb3/sb3-ppo-pendulum.py
Normal file
6
scratch/etienne/trpo/sb3/sb3-ppo-pendulum.py
Normal file
@@ -0,0 +1,6 @@
|
||||
from stable_baselines3 import PPO
|
||||
from stable_baselines3.common.env_util import make_vec_env
|
||||
|
||||
env = make_vec_env("Pendulum-v0", n_envs=4)
|
||||
model = PPO("MlpPolicy", env, verbose=1)
|
||||
model.learn(total_timesteps=250000)
|
||||
16
scratch/etienne/trpo/sb3/sb3-trpo-intersimple.py
Normal file
16
scratch/etienne/trpo/sb3/sb3-trpo-intersimple.py
Normal file
@@ -0,0 +1,16 @@
|
||||
from sb3_contrib import TRPO
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
|
||||
env = IntersimpleLidarFlat(
|
||||
n_rays=5,
|
||||
agent=51,
|
||||
reward=functools.partial(
|
||||
speed_reward,
|
||||
collision_penalty=0
|
||||
),
|
||||
)
|
||||
|
||||
model = TRPO("MlpPolicy", env, use_sde=False, sde_sample_freq=4, verbose=1)
|
||||
model.learn(total_timesteps=250000)
|
||||
8
scratch/etienne/trpo/sb3/sb3-trpo-pendulum.py
Normal file
8
scratch/etienne/trpo/sb3/sb3-trpo-pendulum.py
Normal file
@@ -0,0 +1,8 @@
|
||||
from sb3_contrib import TRPO
|
||||
import gym
|
||||
from gym.wrappers import TransformObservation
|
||||
|
||||
env = TransformObservation(gym.make('Pendulum-v0'), lambda obs: obs)
|
||||
|
||||
model = TRPO("MlpPolicy", env, verbose=1)
|
||||
model.learn(total_timesteps=250000)
|
||||
Reference in New Issue
Block a user