diff --git a/scratch/etienne/intersimple/airl_flat.py b/scratch/etienne/intersimple/airl_flat.py new file mode 100644 index 0000000..e21a38a --- /dev/null +++ b/scratch/etienne/intersimple/airl_flat.py @@ -0,0 +1,64 @@ +# %% +import pathlib +import pickle +import tempfile + +import stable_baselines3 as sb3 +from stable_baselines3.common.env_util import make_vec_env + +from imitation.algorithms import adversarial, bc +from imitation.data import rollout +from imitation.util import logger + +from intersim.envs.intersimple import IntersimpleReward + +model_name = 'airl_flat' + +# Load pickled test demonstrations. +with open("data/NormalizedIntersimpleExpert_IntersimpleRewardAgent51.pkl", "rb") as f: + # This is a list of `imitation.data.types.Trajectory`, where + # every instance contains observations and actions for a single expert + # demonstration. + trajectories = pickle.load(f) + +# %% +# Convert List[types.Trajectory] to an instance of `imitation.data.types.Transitions`. +# This is a more general dataclass containing unordered +# (observation, actions, next_observation) transitions. +transitions = rollout.flatten_trajectories(trajectories) + +venv = make_vec_env(IntersimpleReward, n_envs=2, env_kwargs={'agent': 51}) + +tempdir = tempfile.TemporaryDirectory(prefix="quickstart") +tempdir_path = pathlib.Path(tempdir.name) +print(f"All Tensorboards and logging are being written inside {tempdir_path}/.") + +# Train AIRL on expert data. +# GAIL, and AIRL also accept as `expert_data` any Pytorch-style DataLoader that +# iterates over dictionaries containing observations, actions, and next_observations. +logger.configure(tempdir_path / "AIRL/") +airl_trainer = adversarial.AIRL( + venv, + expert_data=transitions, + expert_batch_size=64, + gen_algo=sb3.PPO("MlpPolicy", venv, verbose=1, n_steps=1024), # n_steps = 2048 ? +) +airl_trainer.train(total_timesteps=100000) +airl_trainer.gen_algo.save(model_name) + +del airl_trainer + +# %% +model = sb3.PPO.load(model_name) + +env = IntersimpleReward(agent=51) + +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) \ No newline at end of file diff --git a/scratch/etienne/intersimple/bc_flat b/scratch/etienne/intersimple/bc_flat new file mode 100644 index 0000000..148f405 Binary files /dev/null and b/scratch/etienne/intersimple/bc_flat differ diff --git a/scratch/etienne/intersimple/bc_flat.py b/scratch/etienne/intersimple/bc_flat.py new file mode 100644 index 0000000..2011a06 --- /dev/null +++ b/scratch/etienne/intersimple/bc_flat.py @@ -0,0 +1,59 @@ +# %% +import pathlib +import pickle +import tempfile + +import stable_baselines3 as sb3 +from stable_baselines3.common.env_util import make_vec_env + +from imitation.algorithms import adversarial, bc +from imitation.data import rollout +from imitation.util import logger + +from intersim.envs.intersimple import IntersimpleReward + +model_name = 'bc_flat' + +# Load pickled test demonstrations. +with open("data/NormalizedIntersimpleExpert_IntersimpleRewardAgent51.pkl", "rb") as f: + # This is a list of `imitation.data.types.Trajectory`, where + # every instance contains observations and actions for a single expert + # demonstration. + trajectories = pickle.load(f) + +# %% +# Convert List[types.Trajectory] to an instance of `imitation.data.types.Transitions`. +# This is a more general dataclass containing unordered +# (observation, actions, next_observation) transitions. +transitions = rollout.flatten_trajectories(trajectories) + +venv = make_vec_env(IntersimpleReward, n_envs=2, env_kwargs={'agent': 51}) + +tempdir = tempfile.TemporaryDirectory(prefix="quickstart") +tempdir_path = pathlib.Path(tempdir.name) +print(f"All Tensorboards and logging are being written inside {tempdir_path}/.") + +# Train BC on expert data. +# BC also accepts as `expert_data` any PyTorch-style DataLoader that iterates over +# dictionaries containing observations and actions. +logger.configure(tempdir_path / "BC/") +bc_trainer = bc.BC(venv.observation_space, venv.action_space, expert_data=transitions) +bc_trainer.train(n_epochs=1000) +bc_trainer.save_policy(model_name) + +del bc_trainer + +# %% +model = bc.reconstruct_policy(model_name) + +env = IntersimpleReward(agent=51) + +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) diff --git a/scratch/etienne/intersimple/data/NormalizedIntersimpleExpert_IntersimpleRewardAgent51.pkl b/scratch/etienne/intersimple/data/NormalizedIntersimpleExpert_IntersimpleRewardAgent51.pkl new file mode 100644 index 0000000..27be74d Binary files /dev/null and b/scratch/etienne/intersimple/data/NormalizedIntersimpleExpert_IntersimpleRewardAgent51.pkl differ diff --git a/scratch/etienne/intersimple/data/NormalizedIntersimpleExpert_NRasterizedAgent51.pkl b/scratch/etienne/intersimple/data/NormalizedIntersimpleExpert_NRasterizedAgent51.pkl new file mode 100644 index 0000000..f875544 Binary files /dev/null and b/scratch/etienne/intersimple/data/NormalizedIntersimpleExpert_NRasterizedAgent51.pkl differ diff --git a/scratch/etienne/intersimple/data/expert.py b/scratch/etienne/intersimple/data/expert.py new file mode 100644 index 0000000..e844cd9 --- /dev/null +++ b/scratch/etienne/intersimple/data/expert.py @@ -0,0 +1,108 @@ +from intersim.envs.intersimple import Intersimple +from stable_baselines3.common.policies import BasePolicy +import gym +import intersim.envs.intersimple +import imitation.data.rollout as rollout +from stable_baselines3.common.vec_env.dummy_vec_env import DummyVecEnv +from imitation.data.wrappers import RolloutInfoWrapper + +class IntersimExpert(BasePolicy): + + def __init__(self, intersim_env, mu=0, *args, **kwargs): + super().__init__( + observation_space=gym.spaces.Space(), + action_space=gym.spaces.Space(), + *args, **kwargs + ) + self._intersim = intersim_env + self._mu = mu + + def forward(self, *args, **kwargs): + raise NotImplementedError() + + def _predict(self, *args, **kwargs): + raise NotImplementedError() + + def _action(self): + target_t = min(self._intersim._ind + 1, len(self._intersim._svt.simstate) - 1) + target_state = self._intersim._svt.simstate[target_t] + return self._intersim.target_state(target_state, mu=self._mu) + + def predict(self, *args, **kwargs): + return self._action(), None + +class IntersimpleExpert(BasePolicy): + + def __init__(self, intersimple_env, mu=0, *args, **kwargs): + super().__init__( + observation_space=intersimple_env.observation_space, + action_space=intersimple_env.action_space, + *args, **kwargs + ) + self._intersimple = intersimple_env + self._intersim_expert = IntersimExpert(intersimple_env._env, mu=mu) + + def forward(self, *args, **kwargs): + raise NotImplementedError() + + def _predict(self, *args, **kwargs): + raise NotImplementedError() + + def _action(self): + return self._intersim_expert._action()[self._intersimple._agent] + + def predict(self, *args, **kwargs): + return self._action(), None + +class NormalizedIntersimpleExpert(IntersimpleExpert): + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + def predict(self, *args, **kwargs): + action, _ = super().predict(*args, **kwargs) + return self._intersimple._normalize(action), None + +class DummyVecEnvPolicy(): + + def __init__(self, experts): + self._experts = [e() for e in experts] + + def predict(self, *args, **kwargs): + predictions = [e.predict() for e in self._experts] + actions = [p[0] for p in predictions] + states = [p[1] for p in predictions] + return actions, states + +def demonstrations(expert='NormalizedIntersimpleExpert', env='NRasterizedRandomAgent', path=None, min_timesteps=25000, min_episodes=None, env_args={}, policy_args={}): + """Rollout and save expert demos. + + Usage: + python -m intersimple.expert + + """ + Env = intersim.envs.intersimple.__dict__[env] + Expert = globals()[expert] + + env = Env(**env_args) + info_env = RolloutInfoWrapper(env) + venv = DummyVecEnv([lambda: info_env]) + + policy = Expert(env, **policy_args) + venv_policy = DummyVecEnvPolicy([lambda: policy]) + + path = path or (policy.__class__.__name__ + '_' + env.__class__.__name__ + '.pkl') + + rollout.rollout_and_save( + path=path, + policy=venv_policy, + venv=venv, + sample_until=rollout.make_sample_until( + min_timesteps=min_timesteps, + min_episodes=min_episodes, + ) + ) + +if __name__ == '__main__': + import fire + fire.Fire(demonstrations) diff --git a/scratch/etienne/intersimple/gail/discriminator.py b/scratch/etienne/intersimple/gail/discriminator.py new file mode 100644 index 0000000..8f26343 --- /dev/null +++ b/scratch/etienne/intersimple/gail/discriminator.py @@ -0,0 +1,51 @@ +import torch + +class CnnDiscriminator(torch.nn.Module): + """ConvNet similar to stable_baselines3.common.policies.ActorCriticCnnPolicy.""" + + def __init__(self, env): + super().__init__() + print('venv obs', env.observation_space.shape) + + obs_channels, _, _ = env.observation_space.shape + (action_size,) = env.action_space.shape + in_channels = obs_channels + action_size + + self.cnn = torch.nn.Sequential( + torch.nn.Conv2d(in_channels, 32, kernel_size=(8, 8), stride=(4, 4)), # 5+1 -> 32 + torch.nn.ReLU(), + torch.nn.Conv2d(32, 64, kernel_size=(4, 4), stride=(2, 2)), # 32 -> 64 + torch.nn.ReLU(), + torch.nn.Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1)), # 64 -> 64 + torch.nn.ReLU(), + torch.nn.Flatten(start_dim=1, end_dim=-1), + torch.nn.LazyLinear(512), # 28224 -> 512 + torch.nn.ReLU(), + torch.nn.LazyLinear(1), # 512 -> 1 + ) + + def forward(self, state, action): + b, _, h, w = state.shape + _, a = action.shape + act_layer = action.unsqueeze(-1).unsqueeze(-1).expand((b, a, h, w)) + sa = torch.cat((act_layer, state), -3) + return self.cnn(sa).squeeze() + +class MlpDiscriminator(torch.nn.Module): + """MLP similar to stable_baselines3.common.policies.ActorCriticPolicy.""" + + def __init__(self, env=None): + super().__init__() + self.flatten = torch.nn.Flatten(start_dim=1, end_dim=-1) + self.mlp = torch.nn.Sequential( + torch.nn.LazyLinear(64), # 42 -> 64 + torch.nn.Tanh(), + torch.nn.LazyLinear(64), # 64 -> 64 + torch.nn.Tanh(), + torch.nn.LazyLinear(1), # 64 -> 1 + ) + + def forward(self, state, action): + flat = self.flatten(state) + sa = torch.cat((action, flat), -1) + return self.mlp(sa).squeeze() diff --git a/scratch/etienne/intersimple/gail_flat.py b/scratch/etienne/intersimple/gail_flat.py new file mode 100644 index 0000000..86af50e --- /dev/null +++ b/scratch/etienne/intersimple/gail_flat.py @@ -0,0 +1,69 @@ +# %% +import pathlib +import pickle +import tempfile + +import stable_baselines3 as sb3 +from stable_baselines3.common.env_util import make_vec_env + +from imitation.algorithms import adversarial, bc +from imitation.data import rollout +from imitation.util import logger + +from intersim.envs.intersimple import IntersimpleReward + +from gail.discriminator import MlpDiscriminator + +model_name = 'gail_flat' + +# %% +# Load pickled test demonstrations. +with open("data/NormalizedIntersimpleExpert_IntersimpleRewardAgent51.pkl", "rb") as f: + # This is a list of `imitation.data.types.Trajectory`, where + # every instance contains observations and actions for a single expert + # demonstration. + trajectories = pickle.load(f) + +# %% +# Convert List[types.Trajectory] to an instance of `imitation.data.types.Transitions`. +# This is a more general dataclass containing unordered +# (observation, actions, next_observation) transitions. +transitions = rollout.flatten_trajectories(trajectories) + +venv = make_vec_env(IntersimpleReward, n_envs=2, env_kwargs={'agent': 51}) + +tempdir = tempfile.TemporaryDirectory(prefix="quickstart") +tempdir_path = pathlib.Path(tempdir.name) +print(f"All Tensorboards and logging are being written inside {tempdir_path}/.") + +# Train GAIL on expert data. +# GAIL, and AIRL also accept as `expert_data` any Pytorch-style DataLoader that +# iterates over dictionaries containing observations, actions, and next_observations. +logger.configure(tempdir_path / "GAIL/") +gail_trainer = adversarial.GAIL( + venv, + expert_data=transitions, + expert_batch_size=220, + #n_disc_updates_per_round=32, + discrim_kwargs={'discrim_net': MlpDiscriminator()}, + gen_algo=sb3.PPO("MlpPolicy", venv, verbose=1, n_steps=4096), +) +gail_trainer.train(total_timesteps=80000) +gail_trainer.gen_algo.save(model_name) + +#del gail_trainer + +# %% +model = sb3.PPO.load(model_name) + +env = IntersimpleReward(agent=51) + +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) diff --git a/scratch/etienne/intersimple/gail_image.py b/scratch/etienne/intersimple/gail_image.py new file mode 100644 index 0000000..1f0ba6a --- /dev/null +++ b/scratch/etienne/intersimple/gail_image.py @@ -0,0 +1,69 @@ +# %% +import pathlib +import pickle +import tempfile + +import stable_baselines3 as sb3 +from stable_baselines3.common.env_util import make_vec_env + +from imitation.algorithms import adversarial, bc +from imitation.data import rollout +from imitation.util import logger + +from intersim.envs.intersimple import NRasterized + +from gail.discriminator import CnnDiscriminator + +model_name = 'gail_image' + +# %% +# Load pickled test demonstrations. +with open("data/NormalizedIntersimpleExpert_NRasterizedAgent51.pkl", "rb") as f: + # This is a list of `imitation.data.types.Trajectory`, where + # every instance contains observations and actions for a single expert + # demonstration. + trajectories = pickle.load(f) + +# %% +# Convert List[types.Trajectory] to an instance of `imitation.data.types.Transitions`. +# This is a more general dataclass containing unordered +# (observation, actions, next_observation) transitions. +transitions = rollout.flatten_trajectories(trajectories) + +venv = make_vec_env(NRasterized, n_envs=2, env_kwargs={'agent': 51}) + +tempdir = tempfile.TemporaryDirectory(prefix="quickstart") +tempdir_path = pathlib.Path(tempdir.name) +print(f"All Tensorboards and logging are being written inside {tempdir_path}/.") + +# Train GAIL on expert data. +# GAIL, and AIRL also accept as `expert_data` any Pytorch-style DataLoader that +# iterates over dictionaries containing observations, actions, and next_observations. +logger.configure(tempdir_path / "GAIL/") +gail_trainer = adversarial.GAIL( + venv, + expert_data=transitions, + expert_batch_size=200, + n_disc_updates_per_round=2048, + discrim_kwargs={'discrim_net': CnnDiscriminator(venv)}, + gen_algo=sb3.PPO("CnnPolicy", venv, verbose=1, n_steps=128), +) +gail_trainer.train(total_timesteps=100000) +gail_trainer.gen_algo.save(model_name) + +#del gail_trainer + +# %% +model = sb3.PPO.load(model_name) + +env = NRasterized(agent=51) + +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) diff --git a/scratch/etienne/intersimple/imitation_quickstart.py b/scratch/etienne/intersimple/imitation_quickstart.py new file mode 100644 index 0000000..996519f --- /dev/null +++ b/scratch/etienne/intersimple/imitation_quickstart.py @@ -0,0 +1,63 @@ +# %% +import pathlib +import pickle +import tempfile + +import stable_baselines3 as sb3 +from stable_baselines3.common.env_util import make_vec_env + +from imitation.algorithms import adversarial, bc +from imitation.data import rollout +from imitation.util import logger + +from intersim.envs.intersimple import IntersimpleReward + +# Load pickled test demonstrations. +with open("data/NormalizedIntersimpleExpert_IntersimpleRewardAgent51.pkl", "rb") as f: + # This is a list of `imitation.data.types.Trajectory`, where + # every instance contains observations and actions for a single expert + # demonstration. + trajectories = pickle.load(f) + +# %% +# Convert List[types.Trajectory] to an instance of `imitation.data.types.Transitions`. +# This is a more general dataclass containing unordered +# (observation, actions, next_observation) transitions. +transitions = rollout.flatten_trajectories(trajectories) + +venv = make_vec_env(IntersimpleReward, n_envs=2, env_kwargs={'agent': 51}) + +tempdir = tempfile.TemporaryDirectory(prefix="quickstart") +tempdir_path = pathlib.Path(tempdir.name) +print(f"All Tensorboards and logging are being written inside {tempdir_path}/.") + +# Train BC on expert data. +# BC also accepts as `expert_data` any PyTorch-style DataLoader that iterates over +# dictionaries containing observations and actions. +logger.configure(tempdir_path / "BC/") +bc_trainer = bc.BC(venv.observation_space, venv.action_space, expert_data=transitions) +bc_trainer.train(n_epochs=1) + +# Train GAIL on expert data. +# GAIL, and AIRL also accept as `expert_data` any Pytorch-style DataLoader that +# iterates over dictionaries containing observations, actions, and next_observations. +logger.configure(tempdir_path / "GAIL/") +gail_trainer = adversarial.GAIL( + venv, + expert_data=transitions, + expert_batch_size=32, + gen_algo=sb3.PPO("MlpPolicy", venv, verbose=1, n_steps=1024), +) +gail_trainer.train(total_timesteps=2048) + +# Train AIRL on expert data. +logger.configure(tempdir_path / "AIRL/") +airl_trainer = adversarial.AIRL( + venv, + expert_data=transitions, + expert_batch_size=32, + gen_algo=sb3.PPO("MlpPolicy", venv, verbose=1, n_steps=1024), +) +airl_trainer.train(total_timesteps=2048) + +# %% diff --git a/scratch/etienne/intersimple/ppo_const.py b/scratch/etienne/intersimple/ppo_const.py new file mode 100644 index 0000000..7fb598f --- /dev/null +++ b/scratch/etienne/intersimple/ppo_const.py @@ -0,0 +1,36 @@ +# %% +from stable_baselines3 import PPO +from stable_baselines3.common.env_util import make_vec_env +from intersim.envs.intersimple import IntersimpleReward, speed_reward + +model_name = "ppo_const" + +env = IntersimpleReward( + agent=51, + #reward=speed_reward, +) + +# %% +model = PPO( + "MlpPolicy", env, + verbose=1, +) +model.learn(total_timesteps=100000) +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) diff --git a/scratch/etienne/intersimple/ppo_const_collision.py b/scratch/etienne/intersimple/ppo_const_collision.py new file mode 100644 index 0000000..ab535c8 --- /dev/null +++ b/scratch/etienne/intersimple/ppo_const_collision.py @@ -0,0 +1,41 @@ +# %% +from stable_baselines3 import PPO +from stable_baselines3.common.env_util import make_vec_env +from intersim.envs.intersimple import ConstCollisionReward, IntersimpleFlatAgent + +model_name = "ppo_const_collision" + +class IntersimpleConstCollisionAgent(ConstCollisionReward, IntersimpleFlatAgent): + pass + +env = IntersimpleConstCollisionAgent( + agent=51, + speed_reward_weight=0.001, + collision_penalty=1000 +) + +# %% +model = PPO( + "MlpPolicy", env, + learning_rate=3e-6, + verbose=1, +) +model.learn(total_timesteps=2e5) +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) diff --git a/scratch/etienne/intersimple/ppo_const_image.py b/scratch/etienne/intersimple/ppo_const_image.py new file mode 100644 index 0000000..eb05b46 --- /dev/null +++ b/scratch/etienne/intersimple/ppo_const_image.py @@ -0,0 +1,35 @@ +# %% +from stable_baselines3 import PPO +from stable_baselines3.common.env_util import make_vec_env +from intersim.envs.intersimple import NRasterized + +model_name = "ppo_const_image" + +env = NRasterized( + agent=51, +) + +# %% +model = PPO( + "CnnPolicy", env, + verbose=1, +) +model.learn(total_timesteps=100000) +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) \ No newline at end of file diff --git a/scratch/etienne/intersimple/ppo_const_image_random.py b/scratch/etienne/intersimple/ppo_const_image_random.py new file mode 100644 index 0000000..da99f18 --- /dev/null +++ b/scratch/etienne/intersimple/ppo_const_image_random.py @@ -0,0 +1,33 @@ +# %% +from stable_baselines3 import PPO +from intersim.envs.intersimple import NRasterizedRandomAgent +import functools + +model_name = "ppo_const_image_random" + +env = NRasterizedRandomAgent() + +# %% +model = PPO( + "CnnPolicy", env, + verbose=1, +) +model.learn(total_timesteps=2e5) +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) diff --git a/scratch/etienne/intersimple/ppo_intersimple_tspeed.py b/scratch/etienne/intersimple/ppo_intersimple_tspeed.py new file mode 100644 index 0000000..3d915f5 --- /dev/null +++ b/scratch/etienne/intersimple/ppo_intersimple_tspeed.py @@ -0,0 +1,25 @@ +from stable_baselines3 import PPO +from stable_baselines3.common.env_util import make_vec_env +from intersim.envs.intersimple import IntersimpleTargetSpeed + +env = IntersimpleTargetSpeed() + +model = PPO("MlpPolicy", env, verbose=1) +model.learn(total_timesteps=25000) +model.save("ppo_intersimple") + +print('Done training.') + +del model # remove to demonstrate saving and loading + +model = PPO.load("ppo_intersimple") + +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() diff --git a/scratch/etienne/intersimple/ppo_speed.py b/scratch/etienne/intersimple/ppo_speed.py new file mode 100644 index 0000000..1aa911b --- /dev/null +++ b/scratch/etienne/intersimple/ppo_speed.py @@ -0,0 +1,46 @@ +# %% +from stable_baselines3 import PPO +from intersim.envs.intersimple import IntersimpleReward, speed_reward +import functools + +model_name = "ppo_speed" + +#def reward(state, action, info): +# speed = state[2].item() +# r = speed if speed < 10 else (10 - 5 * (speed - 10)) +# return 0.1 * r + +env = IntersimpleReward( + agent=51, + reward=functools.partial( + speed_reward, + collision_penalty=0 + ), +) + +# %% +model = PPO( + "MlpPolicy", env, + verbose=1, +) +model.learn(total_timesteps=100000) +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) + +# %% diff --git a/scratch/etienne/intersimple/ppo_speed_image.py b/scratch/etienne/intersimple/ppo_speed_image.py new file mode 100644 index 0000000..2ac95d8 --- /dev/null +++ b/scratch/etienne/intersimple/ppo_speed_image.py @@ -0,0 +1,46 @@ +# %% +from stable_baselines3 import PPO +from intersim.envs.intersimple import NRasterized, speed_reward +import functools + +model_name = "ppo_speed_image" + +#def reward(state, action, info): +# speed = state[2].item() +# r = speed if speed < 10 else (10 - 5 * (speed - 10)) +# return 0.1 * r + +env = NRasterized( + agent=20, + reward=functools.partial( + speed_reward, + collision_penalty=0 + ), +) + +# %% +model = PPO( + "CnnPolicy", env, + verbose=1, +) +model.learn(total_timesteps=100000) +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) + +# %% diff --git a/scratch/etienne/intersimple/ppo_speed_image_random.py b/scratch/etienne/intersimple/ppo_speed_image_random.py new file mode 100644 index 0000000..26fe9e0 --- /dev/null +++ b/scratch/etienne/intersimple/ppo_speed_image_random.py @@ -0,0 +1,39 @@ +# %% +from stable_baselines3 import PPO +from intersim.envs.intersimple import NRasterizedRandomAgent, speed_reward +import functools + +model_name = "ppo_speed_image_random" + +env = NRasterizedRandomAgent( + reward=functools.partial( + speed_reward, + collision_penalty=0 + ) +) + +# %% +model = PPO( + "CnnPolicy", env, + verbose=1, + batch_size=2048, +) +model.learn(total_timesteps=2e5) +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) diff --git a/scratch/etienne/intersimple/ppo_speed_random.py b/scratch/etienne/intersimple/ppo_speed_random.py new file mode 100644 index 0000000..1c5b55f --- /dev/null +++ b/scratch/etienne/intersimple/ppo_speed_random.py @@ -0,0 +1,43 @@ +# %% +from stable_baselines3 import PPO +from intersim.envs.intersimple import IntersimpleFlatRandomAgent, Reward, RewardVisualization, speed_reward +import functools + +model_name = "ppo_speed_random" + +class IntersimpleRewardRandom(RewardVisualization, Reward, IntersimpleFlatRandomAgent): + """`IntersimpleFlatAgent` with rewards.""" + pass + +env = IntersimpleRewardRandom( + reward=functools.partial( + speed_reward, + collision_penalty=0 + ) +) + +# %% +model = PPO( + "MlpPolicy", env, + verbose=1, + batch_size=2048, +) +model.learn(total_timesteps=2e5) +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) diff --git a/scratch/etienne/intersimple/ppo_tspeed.py b/scratch/etienne/intersimple/ppo_tspeed.py new file mode 100644 index 0000000..2d8fc04 --- /dev/null +++ b/scratch/etienne/intersimple/ppo_tspeed.py @@ -0,0 +1,39 @@ +# %% +from stable_baselines3 import PPO +from stable_baselines3.common.env_util import make_vec_env +from intersim.envs.intersimple import IntersimpleTargetSpeedAgent + +model_name = "ppo_tspeed" + +env = IntersimpleTargetSpeedAgent( + agent=51, + target_speed=10, + speed_penalty_weight=0.001, + collision_penalty=1000 +) + +# %% +model = PPO( + "MlpPolicy", env, + learning_rate=3e-6, + verbose=1, +) +model.learn(total_timesteps=2e5) +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) diff --git a/scratch/etienne/intersimple/ppo_tspeed_random.py b/scratch/etienne/intersimple/ppo_tspeed_random.py new file mode 100644 index 0000000..46c9843 --- /dev/null +++ b/scratch/etienne/intersimple/ppo_tspeed_random.py @@ -0,0 +1,31 @@ +# %% +from stable_baselines3 import PPO +from stable_baselines3.common.env_util import make_vec_env +from intersim.envs.intersimple import IntersimpleTargetSpeedRandom + +model_name = "ppo_tspeed_random" + +# %% +env = IntersimpleTargetSpeedRandom(target_speed=10) + +# %% +model = PPO("MlpPolicy", env, verbose=1) +model.learn(total_timesteps=250000) +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()