Copy over experiments
This commit is contained in:
64
scratch/etienne/intersimple/airl_flat.py
Normal file
64
scratch/etienne/intersimple/airl_flat.py
Normal file
@@ -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)
|
||||||
BIN
scratch/etienne/intersimple/bc_flat
Normal file
BIN
scratch/etienne/intersimple/bc_flat
Normal file
Binary file not shown.
59
scratch/etienne/intersimple/bc_flat.py
Normal file
59
scratch/etienne/intersimple/bc_flat.py
Normal file
@@ -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)
|
||||||
Binary file not shown.
Binary file not shown.
108
scratch/etienne/intersimple/data/expert.py
Normal file
108
scratch/etienne/intersimple/data/expert.py
Normal file
@@ -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 <flags>
|
||||||
|
|
||||||
|
"""
|
||||||
|
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)
|
||||||
51
scratch/etienne/intersimple/gail/discriminator.py
Normal file
51
scratch/etienne/intersimple/gail/discriminator.py
Normal file
@@ -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()
|
||||||
69
scratch/etienne/intersimple/gail_flat.py
Normal file
69
scratch/etienne/intersimple/gail_flat.py
Normal file
@@ -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)
|
||||||
69
scratch/etienne/intersimple/gail_image.py
Normal file
69
scratch/etienne/intersimple/gail_image.py
Normal file
@@ -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)
|
||||||
63
scratch/etienne/intersimple/imitation_quickstart.py
Normal file
63
scratch/etienne/intersimple/imitation_quickstart.py
Normal file
@@ -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)
|
||||||
|
|
||||||
|
# %%
|
||||||
36
scratch/etienne/intersimple/ppo_const.py
Normal file
36
scratch/etienne/intersimple/ppo_const.py
Normal file
@@ -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)
|
||||||
41
scratch/etienne/intersimple/ppo_const_collision.py
Normal file
41
scratch/etienne/intersimple/ppo_const_collision.py
Normal file
@@ -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)
|
||||||
35
scratch/etienne/intersimple/ppo_const_image.py
Normal file
35
scratch/etienne/intersimple/ppo_const_image.py
Normal file
@@ -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)
|
||||||
33
scratch/etienne/intersimple/ppo_const_image_random.py
Normal file
33
scratch/etienne/intersimple/ppo_const_image_random.py
Normal file
@@ -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)
|
||||||
25
scratch/etienne/intersimple/ppo_intersimple_tspeed.py
Normal file
25
scratch/etienne/intersimple/ppo_intersimple_tspeed.py
Normal file
@@ -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()
|
||||||
46
scratch/etienne/intersimple/ppo_speed.py
Normal file
46
scratch/etienne/intersimple/ppo_speed.py
Normal file
@@ -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)
|
||||||
|
|
||||||
|
# %%
|
||||||
46
scratch/etienne/intersimple/ppo_speed_image.py
Normal file
46
scratch/etienne/intersimple/ppo_speed_image.py
Normal file
@@ -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)
|
||||||
|
|
||||||
|
# %%
|
||||||
39
scratch/etienne/intersimple/ppo_speed_image_random.py
Normal file
39
scratch/etienne/intersimple/ppo_speed_image_random.py
Normal file
@@ -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)
|
||||||
43
scratch/etienne/intersimple/ppo_speed_random.py
Normal file
43
scratch/etienne/intersimple/ppo_speed_random.py
Normal file
@@ -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)
|
||||||
39
scratch/etienne/intersimple/ppo_tspeed.py
Normal file
39
scratch/etienne/intersimple/ppo_tspeed.py
Normal file
@@ -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)
|
||||||
31
scratch/etienne/intersimple/ppo_tspeed_random.py
Normal file
31
scratch/etienne/intersimple/ppo_tspeed_random.py
Normal file
@@ -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()
|
||||||
Reference in New Issue
Block a user