From 5bd8b42d9fa8768ce0ca03ed4638dfe7418095f3 Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Thu, 17 Feb 2022 22:41:55 +0100 Subject: [PATCH] Merge updated files --- scratch/etienne/trpo/.gitignore | 232 ------------ scratch/etienne/trpo/core/discriminator.py | 74 ---- scratch/etienne/trpo/core/gail.py | 125 ------- scratch/etienne/trpo/core/optimization.py | 39 -- scratch/etienne/trpo/core/policy.py | 96 ----- scratch/etienne/trpo/core/ppo.py | 72 ---- scratch/etienne/trpo/core/reparam_module.py | 162 -------- scratch/etienne/trpo/core/sampling.py | 73 ---- .../etienne/trpo/core/test_optimization.py | 23 -- scratch/etienne/trpo/core/trpo.py | 79 ---- scratch/etienne/trpo/core/value.py | 48 --- scratch/etienne/trpo/core/value_estimation.py | 40 -- .../bc-intersimple-setobs2.py | 2 +- .../gail-intersimple-minobs.py | 2 +- .../gail-intersimple-minobs2.py | 2 +- .../gail-intersimple-normobs.py | 2 +- .../gail-intersimple-setobs.py | 2 +- .../gail-intersimple-setobs2-recurrent.py | 2 +- .../gail-intersimple-setobs2.py | 2 +- .../{ => experiments}/gail-intersimple.py | 2 +- .../{ => experiments}/gail-options-minobs.py | 2 +- .../{ => experiments}/gail-options-setobs.py | 2 +- .../{ => experiments}/gail-options-setobs2.py | 10 +- .../trpo/{ => experiments}/gail-pendulum.py | 0 .../gail-ppo-intersimple-minobs.py | 2 +- .../gail-ppo-intersimple-normobs.py | 2 +- .../gail-ppo-intersimple-setobs2.py | 2 +- .../{ => experiments}/gail-ppo-intersimple.py | 2 +- .../gail-ppo-options-minobs.py | 2 +- .../gail-ppo-options-setobs.py | 2 +- .../gail-ppo-options-setobs2.py | 8 +- .../intersimple-expert-action-profiles.ipynb | 0 .../intersimple-expert-rollout-minobs.py | 2 +- .../intersimple-expert-rollout-minobs2.py | 2 +- .../intersimple-expert-rollout-normobs.py | 2 +- .../intersimple-expert-rollout-setobs.py | 2 +- .../intersimple-expert-rollout-setobs2.py | 2 +- .../intersimple-expert-rollout.py | 2 +- .../ppo-intersimple-minobs.py | 2 +- .../ppo-intersimple-minobs2.py | 2 +- .../ppo-intersimple-normobs.py | 0 .../trpo/{ => experiments}/ppo-intersimple.py | 0 .../{ => experiments}/ppo-options-minobs.py | 4 +- .../trpo/{ => experiments}/ppo-pendulum.py | 0 .../etienne/trpo/{ => experiments}/readme.md | 0 .../etienne/trpo/experiments/requirements.txt | 3 + .../trpo/experiments/sgail-options-setobs2.py | 108 ++++++ .../experiments/sgail-ppo-options-setobs2.py | 107 ++++++ .../trpo-intersimple-minobs.py | 4 +- .../trpo-intersimple-minobs2.py | 4 +- .../trpo-intersimple-normobs.py | 0 .../trpo-intersimple-setobs.py | 4 +- .../trpo-intersimple-setobs2.py | 4 +- .../{ => experiments}/trpo-intersimple.py | 0 .../{ => experiments}/trpo-options-minobs.py | 4 +- .../trpo-pendulum-rollout.py | 0 .../trpo/{ => experiments}/trpo-pendulum.py | 0 .../trpo/{ => experiments}/trpo-walker.py | 0 .../etienne/trpo/experiments/vec-env.ipynb | 346 ++++++++++++++++++ .../wgail-intersimple-minobs.py | 2 +- .../wgail-intersimple-minobs2.py | 2 +- .../wgail-intersimple-setobs2.py | 2 +- .../{ => experiments}/wgail-intersimple.py | 2 +- .../{ => experiments}/wgail-options-setobs.py | 2 +- .../wgail-options-setobs2.py | 4 +- .../trpo/{ => experiments}/wgail-pendulum.py | 0 .../wgail-ppo-intersimple-minobs.py | 2 +- .../wgail-ppo-intersimple-setobs2.py | 2 +- .../wgail-ppo-intersimple.py | 0 .../wgail-ppo-options-setobs.py | 2 +- .../wgail-ppo-options-setobs2.py | 6 +- .../{ => experiments}/wgail-ppo-pendulum.py | 0 .../trpo/sb3/sb3-ppo-intersimple-rollout.py | 2 +- src/core/reparam_module.py | 3 + src/options/envs.py | 111 ++++++ .../etienne/trpo => src}/options/options.py | 17 +- .../trpo => src}/options/test_options.py | 0 src/safe_options/collisions.py | 185 ++++++++++ src/safe_options/options.py | 305 +++++++++++++++ src/safe_options/policy.py | 35 ++ src/safe_options/policy_gradient.py | 107 ++++++ src/safe_options/test_options.py | 54 +++ .../etienne/trpo => src/util}/wrappers.py | 0 83 files changed, 1439 insertions(+), 1123 deletions(-) delete mode 100644 scratch/etienne/trpo/.gitignore delete mode 100644 scratch/etienne/trpo/core/discriminator.py delete mode 100644 scratch/etienne/trpo/core/gail.py delete mode 100644 scratch/etienne/trpo/core/optimization.py delete mode 100644 scratch/etienne/trpo/core/policy.py delete mode 100644 scratch/etienne/trpo/core/ppo.py delete mode 100644 scratch/etienne/trpo/core/reparam_module.py delete mode 100644 scratch/etienne/trpo/core/sampling.py delete mode 100644 scratch/etienne/trpo/core/test_optimization.py delete mode 100644 scratch/etienne/trpo/core/trpo.py delete mode 100644 scratch/etienne/trpo/core/value.py delete mode 100644 scratch/etienne/trpo/core/value_estimation.py rename scratch/etienne/trpo/{ => experiments}/bc-intersimple-setobs2.py (96%) rename scratch/etienne/trpo/{ => experiments}/gail-intersimple-minobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-intersimple-minobs2.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-intersimple-normobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-intersimple-setobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-intersimple-setobs2-recurrent.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-intersimple-setobs2.py (98%) rename scratch/etienne/trpo/{ => experiments}/gail-intersimple.py (96%) rename scratch/etienne/trpo/{ => experiments}/gail-options-minobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-options-setobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-options-setobs2.py (89%) rename scratch/etienne/trpo/{ => experiments}/gail-pendulum.py (100%) rename scratch/etienne/trpo/{ => experiments}/gail-ppo-intersimple-minobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-ppo-intersimple-normobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-ppo-intersimple-setobs2.py (98%) rename scratch/etienne/trpo/{ => experiments}/gail-ppo-intersimple.py (96%) rename scratch/etienne/trpo/{ => experiments}/gail-ppo-options-minobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-ppo-options-setobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-ppo-options-setobs2.py (89%) rename scratch/etienne/trpo/{ => experiments}/intersimple-expert-action-profiles.ipynb (100%) rename scratch/etienne/trpo/{ => experiments}/intersimple-expert-rollout-minobs.py (96%) rename scratch/etienne/trpo/{ => experiments}/intersimple-expert-rollout-minobs2.py (96%) rename scratch/etienne/trpo/{ => experiments}/intersimple-expert-rollout-normobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/intersimple-expert-rollout-setobs.py (96%) rename scratch/etienne/trpo/{ => experiments}/intersimple-expert-rollout-setobs2.py (96%) rename scratch/etienne/trpo/{ => experiments}/intersimple-expert-rollout.py (94%) rename scratch/etienne/trpo/{ => experiments}/ppo-intersimple-minobs.py (98%) rename scratch/etienne/trpo/{ => experiments}/ppo-intersimple-minobs2.py (98%) rename scratch/etienne/trpo/{ => experiments}/ppo-intersimple-normobs.py (100%) rename scratch/etienne/trpo/{ => experiments}/ppo-intersimple.py (100%) rename scratch/etienne/trpo/{ => experiments}/ppo-options-minobs.py (95%) rename scratch/etienne/trpo/{ => experiments}/ppo-pendulum.py (100%) rename scratch/etienne/trpo/{ => experiments}/readme.md (100%) create mode 100644 scratch/etienne/trpo/experiments/requirements.txt create mode 100644 scratch/etienne/trpo/experiments/sgail-options-setobs2.py create mode 100644 scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py rename scratch/etienne/trpo/{ => experiments}/trpo-intersimple-minobs.py (96%) rename scratch/etienne/trpo/{ => experiments}/trpo-intersimple-minobs2.py (96%) rename scratch/etienne/trpo/{ => experiments}/trpo-intersimple-normobs.py (100%) rename scratch/etienne/trpo/{ => experiments}/trpo-intersimple-setobs.py (96%) rename scratch/etienne/trpo/{ => experiments}/trpo-intersimple-setobs2.py (96%) rename scratch/etienne/trpo/{ => experiments}/trpo-intersimple.py (100%) rename scratch/etienne/trpo/{ => experiments}/trpo-options-minobs.py (95%) rename scratch/etienne/trpo/{ => experiments}/trpo-pendulum-rollout.py (100%) rename scratch/etienne/trpo/{ => experiments}/trpo-pendulum.py (100%) rename scratch/etienne/trpo/{ => experiments}/trpo-walker.py (100%) create mode 100644 scratch/etienne/trpo/experiments/vec-env.ipynb rename scratch/etienne/trpo/{ => experiments}/wgail-intersimple-minobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/wgail-intersimple-minobs2.py (97%) rename scratch/etienne/trpo/{ => experiments}/wgail-intersimple-setobs2.py (97%) rename scratch/etienne/trpo/{ => experiments}/wgail-intersimple.py (96%) rename scratch/etienne/trpo/{ => experiments}/wgail-options-setobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/wgail-options-setobs2.py (96%) rename scratch/etienne/trpo/{ => experiments}/wgail-pendulum.py (100%) rename scratch/etienne/trpo/{ => experiments}/wgail-ppo-intersimple-minobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/wgail-ppo-intersimple-setobs2.py (97%) rename scratch/etienne/trpo/{ => experiments}/wgail-ppo-intersimple.py (100%) rename scratch/etienne/trpo/{ => experiments}/wgail-ppo-options-setobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/wgail-ppo-options-setobs2.py (96%) rename scratch/etienne/trpo/{ => experiments}/wgail-ppo-pendulum.py (100%) create mode 100644 src/options/envs.py rename {scratch/etienne/trpo => src}/options/options.py (96%) rename {scratch/etienne/trpo => src}/options/test_options.py (100%) create mode 100644 src/safe_options/collisions.py create mode 100644 src/safe_options/options.py create mode 100644 src/safe_options/policy.py create mode 100644 src/safe_options/policy_gradient.py create mode 100644 src/safe_options/test_options.py rename {scratch/etienne/trpo => src/util}/wrappers.py (100%) diff --git a/scratch/etienne/trpo/.gitignore b/scratch/etienne/trpo/.gitignore deleted file mode 100644 index 7491629..0000000 --- a/scratch/etienne/trpo/.gitignore +++ /dev/null @@ -1,232 +0,0 @@ -PyTorch-Reparam-Module -cg.ipynb -vec-env.ipynb -*.zip -*.pt -*.mp4 -*.pkl -runs/ - -# Created by https://www.toptal.com/developers/gitignore/api/linux,macos,python,visualstudiocode -# Edit at https://www.toptal.com/developers/gitignore?templates=linux,macos,python,visualstudiocode - -### Linux ### -*~ - -# temporary files which can be created if a process still has a handle open of a deleted file -.fuse_hidden* - -# KDE directory preferences -.directory - -# Linux trash folder which might appear on any partition or disk -.Trash-* - -# .nfs files are created when an open file is removed but is still being accessed -.nfs* - -### macOS ### -# General -.DS_Store -.AppleDouble -.LSOverride - -# Icon must end with two \r -Icon - - -# Thumbnails -._* - -# Files that might appear in the root of a volume -.DocumentRevisions-V100 -.fseventsd -.Spotlight-V100 -.TemporaryItems -.Trashes -.VolumeIcon.icns -.com.apple.timemachine.donotpresent - -# Directories potentially created on remote AFP share -.AppleDB -.AppleDesktop -Network Trash Folder -Temporary Items -.apdisk - -### Python ### -# Byte-compiled / optimized / DLL files -__pycache__/ -*.py[cod] -*$py.class - -# C extensions -*.so - -# Distribution / packaging -.Python -build/ -develop-eggs/ -dist/ -downloads/ -eggs/ -.eggs/ -lib/ -lib64/ -parts/ -sdist/ -var/ -wheels/ -share/python-wheels/ -*.egg-info/ -.installed.cfg -*.egg -MANIFEST - -# PyInstaller -# Usually these files are written by a python script from a template -# before PyInstaller builds the exe, so as to inject date/other infos into it. -*.manifest -*.spec - -# Installer logs -pip-log.txt -pip-delete-this-directory.txt - -# Unit test / coverage reports -htmlcov/ -.tox/ -.nox/ -.coverage -.coverage.* -.cache -nosetests.xml -coverage.xml -*.cover -*.py,cover -.hypothesis/ -.pytest_cache/ -cover/ - -# Translations -*.mo -*.pot - -# Django stuff: -*.log -local_settings.py -db.sqlite3 -db.sqlite3-journal - -# Flask stuff: -instance/ -.webassets-cache - -# Scrapy stuff: -.scrapy - -# Sphinx documentation -docs/_build/ - -# PyBuilder -.pybuilder/ -target/ - -# Jupyter Notebook -.ipynb_checkpoints - -# IPython -profile_default/ -ipython_config.py - -# pyenv -# For a library or package, you might want to ignore these files since the code is -# intended to run in multiple environments; otherwise, check them in: -# .python-version - -# pipenv -# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control. -# However, in case of collaboration, if having platform-specific dependencies or dependencies -# having no cross-platform support, pipenv may install dependencies that don't work, or not -# install all needed dependencies. -#Pipfile.lock - -# poetry -# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control. -# This is especially recommended for binary packages to ensure reproducibility, and is more -# commonly ignored for libraries. -# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control -#poetry.lock - -# PEP 582; used by e.g. github.com/David-OConnor/pyflow -__pypackages__/ - -# Celery stuff -celerybeat-schedule -celerybeat.pid - -# SageMath parsed files -*.sage.py - -# Environments -.env -.venv -env/ -venv/ -ENV/ -env.bak/ -venv.bak/ - -# Spyder project settings -.spyderproject -.spyproject - -# Rope project settings -.ropeproject - -# mkdocs documentation -/site - -# mypy -.mypy_cache/ -.dmypy.json -dmypy.json - -# Pyre type checker -.pyre/ - -# pytype static type analyzer -.pytype/ - -# Cython debug symbols -cython_debug/ - -# PyCharm -# JetBrains specific template is maintainted in a separate JetBrains.gitignore that can -# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore -# and can be added to the global gitignore or merged into this file. For a more nuclear -# option (not recommended) you can uncomment the following to ignore the entire idea folder. -#.idea/ - -### VisualStudioCode ### -.vscode/* -!.vscode/settings.json -!.vscode/tasks.json -!.vscode/launch.json -!.vscode/extensions.json -!.vscode/*.code-snippets - -# Local History for Visual Studio Code -.history/ - -# Built Visual Studio Code Extensions -*.vsix - -### VisualStudioCode Patch ### -# Ignore all local history of files -.history -.ionide - -# Support for Project snippet scope - -# End of https://www.toptal.com/developers/gitignore/api/linux,macos,python,visualstudiocode \ No newline at end of file diff --git a/scratch/etienne/trpo/core/discriminator.py b/scratch/etienne/trpo/core/discriminator.py deleted file mode 100644 index 8074c32..0000000 --- a/scratch/etienne/trpo/core/discriminator.py +++ /dev/null @@ -1,74 +0,0 @@ -import torch -import torch.nn as nn - -class Discriminator(nn.Module): - - def __init__(self): - super().__init__() - self.nn = nn.Sequential( - nn.LazyLinear(50), - nn.Tanh(), - nn.LazyLinear(50), - nn.Tanh(), - nn.LazyLinear(1), - ) - - def forward(self, states, actions): - return self.nn(torch.cat((states, actions), dim=-1)).squeeze(-1) - -class DeepsetDiscriminator(nn.Module): - - def __init__(self): - super().__init__() - self.elem = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - ) - self.glob = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(1), - ) - - def forward(self, states, actions): - actions = actions.unsqueeze(-2) - actions = actions.expand(*actions.shape[:-2], states.shape[-2], actions.shape[-1]) - sa = torch.cat((states, actions), dim=-1) - return self.glob(self.elem(sa).sum(-2)).squeeze(-1) - -class RecurrentDiscriminator(nn.Module): - - def __init__(self): - super().__init__() - self.state_dim = 10 - self.state = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(self.state_dim), - ) - self.glob = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(1), - ) - - def forward(self, states, actions): - actions = actions.unsqueeze(-2) - batch_size = actions.shape[:-2] - set_size = states.shape[-2] - action_dim = actions.shape[-1] - actions = actions.expand(*batch_size, set_size, action_dim) - sa = torch.cat((states, actions), dim=-1) - - state = torch.zeros((*batch_size, self.state_dim)) - for i in range(set_size): - state = state + self.state(torch.cat((state, sa[..., i, :]), dim=-1)) - - return self.glob(state).squeeze(-1) diff --git a/scratch/etienne/trpo/core/gail.py b/scratch/etienne/trpo/core/gail.py deleted file mode 100644 index d630bd4..0000000 --- a/scratch/etienne/trpo/core/gail.py +++ /dev/null @@ -1,125 +0,0 @@ -import torch -import torch.nn.functional as F -from dataclasses import dataclass -from core.reparam_module import ReparamPolicy -from core.sampling import rollout -from core.trpo import trpo_step -from core.ppo import ppo_step -from tqdm import tqdm - -class TerminalLogger: - def add_scalar(self, key, scalar, i=None): - if i is not None: - print('Iteration', i, end=' ') - print(key, scalar) - -@dataclass -class Buffer: - states: torch.Tensor - actions: torch.Tensor - rewards: torch.Tensor - dones: torch.Tensor - -def roll_buffer(buffer, *args, **kwargs): - return Buffer( - torch.roll(buffer.states, *args, **kwargs), - torch.roll(buffer.actions, *args, **kwargs), - torch.roll(buffer.rewards, *args, **kwargs), - torch.roll(buffer.dones, *args, **kwargs), - ) - -def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, - v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, - gae_lambda, delta, backtrack_coeff, backtrack_iters, cg_iters=10, cg_damping=0.1, wasserstein=False, wasserstein_c=None, logger=TerminalLogger()): - - policy(torch.zeros(env_fn(0).observation_space.shape)) - policy = ReparamPolicy(policy) - - logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0]) - logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0]) - - for epoch in tqdm(range(epochs)): - generator_data = Buffer(*rollout(env_fn, policy, rollout_episodes, rollout_steps)) - - logger.add_scalar('gen/mean_episode_length', (~generator_data.dones).sum() / generator_data.states.shape[0], epoch) - logger.add_scalar('gen/mean_reward_per_episode', generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0], epoch) - - discriminator, loss = train_discriminator(expert_data, generator_data, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c) - if wasserstein: - generator_data.rewards = discriminator(generator_data.states, generator_data.actions) - else: - generator_data.rewards = -F.logsigmoid(discriminator(generator_data.states, generator_data.actions)) - logger.add_scalar('disc/final_loss', loss, epoch) - logger.add_scalar('disc/mean_reward_per_episode', generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0], epoch) - - value, policy = trpo_step(value, policy, generator_data.states, generator_data.actions, generator_data.rewards, generator_data.dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters, cg_damping) - expert_data = roll_buffer(expert_data, shifts=-3, dims=0) - - return value, policy - -def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, - v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, - gae_lambda, clip_ratio, pi_opt, pi_iters, target_kl=None, max_grad_norm=None, wasserstein=False, wasserstein_c=None, logger=TerminalLogger()): - - logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0]) - logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0]) - - for epoch in range(epochs): - generator_data = Buffer(*rollout(env_fn, policy, rollout_episodes, rollout_steps)) - - logger.add_scalar('gen/mean_episode_length', (~generator_data.dones).sum() / generator_data.states.shape[0], epoch) - logger.add_scalar('gen/mean_reward_per_episode', generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0], epoch) - - discriminator, loss = train_discriminator(expert_data, generator_data, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c) - if wasserstein: - generator_data.rewards = discriminator(generator_data.states, generator_data.actions) - else: - generator_data.rewards = -F.logsigmoid(discriminator(generator_data.states, generator_data.actions)) - logger.add_scalar('disc/final_loss', loss, epoch) - logger.add_scalar('disc/mean_reward_per_episode', generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0], epoch) - - value, policy = ppo_step(value, policy, generator_data.states, generator_data.actions, generator_data.rewards, generator_data.dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm) - expert_data = roll_buffer(expert_data, shifts=-3, dims=0) - - return value, policy - -def train_discriminator(expert_data, generator_data, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c=None): - - n_expert_samples = (~expert_data.dones).sum() - n_generator_samples = (~generator_data.dones).sum() - n_samples = torch.minimum(n_expert_samples, n_generator_samples) - - gen_states = generator_data.states[~generator_data.dones][:n_samples] - gen_actions = generator_data.actions[~generator_data.dones][:n_samples] - exp_states = expert_data.states[~expert_data.dones][:n_samples] - exp_actions = expert_data.actions[~expert_data.dones][:n_samples] - - states = torch.cat((exp_states, gen_states), dim=0).detach() - actions = torch.cat((exp_actions, gen_actions), dim=0).detach() - labels = torch.cat((torch.zeros(n_samples), torch.ones(n_samples))).detach() - - # print('Batch augmentation on') - # random_states = torch.rand_like(gen_states) - # random_actions = torch.rand_like(gen_actions) - # states = torch.cat((exp_states, gen_states, random_states), dim=0).detach() - # actions = torch.cat((exp_actions, gen_actions, random_actions), dim=0).detach() - # labels = torch.cat((torch.zeros(n_samples), torch.ones(n_samples), torch.ones(n_samples))).detach() - - for _ in range(disc_iters): - disc_opt.zero_grad() - pred = discriminator(states, actions) - - if wasserstein: - loss = -(pred * (1 - labels) - pred * labels).mean() - else: - loss = F.binary_cross_entropy(torch.sigmoid(pred), labels) - - loss.backward() - disc_opt.step() - - if wasserstein_c is not None: - with torch.no_grad(): - for param in discriminator.parameters(): - param.clamp_(-wasserstein_c, wasserstein_c) - - return discriminator, loss diff --git a/scratch/etienne/trpo/core/optimization.py b/scratch/etienne/trpo/core/optimization.py deleted file mode 100644 index 4057214..0000000 --- a/scratch/etienne/trpo/core/optimization.py +++ /dev/null @@ -1,39 +0,0 @@ -import torch - -def conjugate_gradient(A, b, max_iters, res_tol=1e-10): - x = torch.zeros_like(b) - r = b - A(x) - p = r - - rTr = r.T @ r - - for _ in range(max_iters): - Ap = A(p) - alpha = rTr / (p.T @ Ap) - x = x + alpha * p - - r = r - alpha * Ap - if torch.norm(r) < res_tol: - break - - rTrnew = r.T @ r - beta = rTrnew / rTr - p = r + beta * p - rTr = rTrnew - - return x - -def line_search(f, x0, dx, g0, alpha, condition, max_steps=10, c1=0.1): - assert 0 < alpha < 1 - - f0 = f(x0) - for _ in range(max_steps): - x = x0 + dx - - if (f(x) > f0 + c1 * g0.T @ dx) and condition(x): - return x - - dx *= alpha - - print('Line search failed, returning x0') - return x0 diff --git a/scratch/etienne/trpo/core/policy.py b/scratch/etienne/trpo/core/policy.py deleted file mode 100644 index 579e72b..0000000 --- a/scratch/etienne/trpo/core/policy.py +++ /dev/null @@ -1,96 +0,0 @@ -import torch -import torch.nn as nn -from torch.distributions import Independent, Normal, Categorical -from torch.distributions.kl import kl_divergence - -class BasePolicy(nn.Module): - - def __init__(self, action_dim): - super().__init__() - self.action_dim = action_dim - - def torch_dist(self, dist): - return Independent(Normal(dist[..., :self.action_dim], dist[..., self.action_dim:].exp()), 1) - - def sample(self, dist): - return self.torch_dist(dist).sample() - - def predict(self, states): - return self.sample(self.forward(states)) - - def log_prob(self, dist, actions): - return self.torch_dist(dist).log_prob(actions) - - def kl_divergence(self, dist1, dist2): - d1 = self.torch_dist(dist1) - d2 = self.torch_dist(dist2) - return kl_divergence(d1, d2) - -class Policy(BasePolicy): - - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - self.nn = nn.Sequential( - nn.LazyLinear(50), - nn.Tanh(), - nn.LazyLinear(50), - nn.Tanh(), - nn.LazyLinear(2 * self.action_dim), - ) - - def forward(self, states): - return self.nn(states) - -class DiscretePolicy(BasePolicy): - - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - self.nn = nn.Sequential( - nn.LazyLinear(50), - nn.Tanh(), - nn.LazyLinear(50), - nn.Tanh(), - nn.LazyLinear(self.action_dim), - ) - - def forward(self, states): - return self.nn(states) - - def torch_dist(self, dist): - return Categorical(logits=dist) - -class SetPolicy(Policy): - - def forward(self, states): - batch_size = states.shape[:-2] - states = torch.cat((states[..., :1, [0, 1]], states[..., :, [2, 5]]), axis=-2).reshape(*batch_size, -1) - return super().forward(states) - -class SetDiscretePolicy(DiscretePolicy): - - def forward(self, states): - batch_size = states.shape[:-2] - states = torch.cat((states[..., :1, [0, 1]], states[..., :, [2, 5]]), axis=-2).reshape(*batch_size, -1) - return super().forward(states) - -class DeepSetPolicy(BasePolicy): - - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - self.elem = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - ) - self.glob = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(2 * self.action_dim), - ) - - def forward(self, states): - return self.glob(self.elem(states).sum(-2)) diff --git a/scratch/etienne/trpo/core/ppo.py b/scratch/etienne/trpo/core/ppo.py deleted file mode 100644 index c742061..0000000 --- a/scratch/etienne/trpo/core/ppo.py +++ /dev/null @@ -1,72 +0,0 @@ -import torch -from core.sampling import rollout -from core.value_estimation import gae - -def ppo(env_fn, value, policy, epochs, rollout_episodes, rollout_steps, gamma, gae_lambda, clip_ratio, pi_opt, pi_iters, v_opt, v_iters, target_kl=None, max_grad_norm=None): - - for epoch in range(epochs): - policy.eval() - states, actions, rewards, dones = rollout(env_fn, policy, rollout_episodes, rollout_steps) - - print('mean', states[~dones].mean(0)) - print('std', states[~dones].std(0)) - - print(f'Iteration {epoch} mean episode length {(~dones).sum() / states.shape[0]}') - print(f'Iteration {epoch} mean reward per episode {rewards[~dones].sum() / states.shape[0]}') - - policy.train() - value.train() - value, policy = ppo_step(value, policy, states, actions, rewards, dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm) - - return value, policy - -def ppo_step(value, policy, states, actions, rewards, dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm): - - states = states.detach() - actions = actions.detach() - rewards = rewards.detach() - dones = dones.detach() - - advantages, returns, valid = gae(states, rewards, value(states), dones, gamma, gae_lambda) - advantages = advantages.detach() - returns = returns.detach() - - # update value function - - for _ in range(v_iters): - v_opt.zero_grad() - value_loss = (value(states) - returns).pow(2)[valid].mean() - value_loss.backward() - v_opt.step() - - # update policy - - old_dist = policy(states).detach() - old_logprob = policy.log_prob(old_dist, actions).detach() - - def g(advantages, clip_ratio): - return torch.where(advantages >= 0, (1 + clip_ratio) * advantages, (1 - clip_ratio) * advantages) - - def L(states, actions, advantages, clip_ratio): - return torch.minimum( - (policy.log_prob(policy(states), actions) - old_logprob).exp() * advantages, - g(advantages, clip_ratio) - )[valid].mean() - - for _ in range(pi_iters): - pi_opt.zero_grad() - ppo_loss = -L(states, actions, advantages, clip_ratio) - ppo_loss.backward() - - if max_grad_norm: - torch.nn.utils.clip_grad_norm(policy.parameters(), max_grad_norm) - - pi_opt.step() - - kl = policy.kl_divergence(policy(states), old_dist)[valid].mean() - if target_kl and kl > target_kl: - break - - print('KL', kl.item()) - - return value, policy diff --git a/scratch/etienne/trpo/core/reparam_module.py b/scratch/etienne/trpo/core/reparam_module.py deleted file mode 100644 index 5bcd613..0000000 --- a/scratch/etienne/trpo/core/reparam_module.py +++ /dev/null @@ -1,162 +0,0 @@ -# Source: https://github.com/SsnL/PyTorch-Reparam-Module - -import torch -import torch.nn as nn -import warnings -import types -from collections import namedtuple -from contextlib import contextmanager - -class ReparamModule(nn.Module): - def __init__(self, module): - super(ReparamModule, self).__init__() - self.module = module - - param_infos = [] - shared_param_memo = {} - shared_param_infos = [] - params = [] - param_numels = [] - param_shapes = [] - for m in self.modules(): - for n, p in m.named_parameters(recurse=False): - if p is not None: - if p in shared_param_memo: - shared_m, shared_n = shared_param_memo[p] - shared_param_infos.append((m, n, shared_m, shared_n)) - else: - shared_param_memo[p] = (m, n) - param_infos.append((m, n)) - params.append(p.detach()) - param_numels.append(p.numel()) - param_shapes.append(p.size()) - - assert len(set(p.dtype for p in params)) <= 1, \ - "expects all parameters in module to have same dtype" - - # store the info for unflatten - self._param_infos = tuple(param_infos) - self._shared_param_infos = tuple(shared_param_infos) - self._param_numels = tuple(param_numels) - self._param_shapes = tuple(param_shapes) - - # flatten - flat_param = nn.Parameter(torch.cat([p.reshape(-1) for p in params], 0)) - self.register_parameter('flat_param', flat_param) - self.param_numel = flat_param.numel() - del params - del shared_param_memo - - # deregister the names as parameters - for m, n in self._param_infos: - delattr(m, n) - for m, n, _, _ in self._shared_param_infos: - delattr(m, n) - - # register the views as plain attributes - self._unflatten_param(self.flat_param) - - # now buffers - # they are not reparametrized. just store info as (module, name, buffer) - buffer_infos = [] - for m in self.modules(): - for n, b in m.named_buffers(recurse=False): - if b is not None: - buffer_infos.append((m, n, b)) - - self._buffer_infos = tuple(buffer_infos) - self._traced_self = None - - def trace(self, example_input, **trace_kwargs): - assert self._traced_self is None, 'This ReparamModule is already traced' - - if isinstance(example_input, torch.Tensor): - example_input = (example_input,) - example_input = tuple(example_input) - example_param = (self.flat_param.detach().clone(),) - example_buffers = (tuple(b.detach().clone() for _, _, b in self._buffer_infos),) - - self._traced_self = torch.jit.trace_module( - self, - inputs=dict( - _forward_with_param=example_param + example_input, - _forward_with_param_and_buffers=example_param + example_buffers + example_input, - ), - **trace_kwargs, - ) - - # replace forwards with traced versions - self._forward_with_param = self._traced_self._forward_with_param - self._forward_with_param_and_buffers = self._traced_self._forward_with_param_and_buffers - return self - - def clear_views(self): - for m, n in self._param_infos: - setattr(m, n, None) # This will set as plain attr - - def _apply(self, *args, **kwargs): - if self._traced_self is not None: - self._traced_self._apply(*args, **kwargs) - return self - return super(ReparamModule, self)._apply(*args, **kwargs) - - def _unflatten_param(self, flat_param): - ps = (t.view(s) for (t, s) in zip(flat_param.split(self._param_numels), self._param_shapes)) - for (m, n), p in zip(self._param_infos, ps): - setattr(m, n, p) # This will set as plain attr - for (m, n, shared_m, shared_n) in self._shared_param_infos: - setattr(m, n, getattr(shared_m, shared_n)) - - @contextmanager - def unflattened_param(self, flat_param): - saved_views = [getattr(m, n) for m, n in self._param_infos] - self._unflatten_param(flat_param) - yield - # Why not just `self._unflatten_param(self.flat_param)`? - # 1. because of https://github.com/pytorch/pytorch/issues/17583 - # 2. slightly faster since it does not require reconstruct the split+view - # graph - for (m, n), p in zip(self._param_infos, saved_views): - setattr(m, n, p) - for (m, n, shared_m, shared_n) in self._shared_param_infos: - setattr(m, n, getattr(shared_m, shared_n)) - - @contextmanager - def replaced_buffers(self, buffers): - for (m, n, _), new_b in zip(self._buffer_infos, buffers): - setattr(m, n, new_b) - yield - for m, n, old_b in self._buffer_infos: - setattr(m, n, old_b) - - def _forward_with_param_and_buffers(self, flat_param, buffers, *inputs, **kwinputs): - with self.unflattened_param(flat_param): - with self.replaced_buffers(buffers): - return self.module(*inputs, **kwinputs) - - def _forward_with_param(self, flat_param, *inputs, **kwinputs): - with self.unflattened_param(flat_param): - return self.module(*inputs, **kwinputs) - - def forward(self, *inputs, flat_param=None, buffers=None, **kwinputs): - if flat_param is None: - flat_param = self.flat_param - if buffers is None: - return self._forward_with_param(flat_param, *inputs, **kwinputs) - else: - return self._forward_with_param_and_buffers(flat_param, tuple(buffers), *inputs, **kwinputs) - - -class ReparamPolicy(ReparamModule): - - def sample(self, *args, **kwargs): - return self.module.sample(*args, **kwargs) - - def log_prob(self, *args, **kwargs): - return self.module.log_prob(*args, **kwargs) - - def kl_divergence(self, *args, **kwargs): - return self.module.kl_divergence(*args, **kwargs) - - def predict(self, *args, **kwargs): - return self.module.predict(*args, **kwargs) diff --git a/scratch/etienne/trpo/core/sampling.py b/scratch/etienne/trpo/core/sampling.py deleted file mode 100644 index 66fbbef..0000000 --- a/scratch/etienne/trpo/core/sampling.py +++ /dev/null @@ -1,73 +0,0 @@ -import torch -import gym -from stable_baselines3.common.vec_env import DummyVecEnv as VecEnv -from tqdm import tqdm - -def rollout(env_fn, policy, n_episodes, max_steps_per_episode): - env = env_fn(0) - states = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.observation_space.shape) - actions = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.action_space.shape) - rewards = torch.zeros(n_episodes, max_steps_per_episode + 1) - dones = torch.ones(n_episodes, max_steps_per_episode + 1, dtype=bool) - - env = VecEnv(list(map(lambda i: (lambda: env_fn(i)), range(n_episodes)))) - - states[:, 0] = torch.tensor(env.reset()).clone().detach() - dones[:, 0] = False - - for s in range(max_steps_per_episode): - actions[:, s] = policy.sample(policy(states[:, s])).clone().detach() - - clipped_actions = actions[:, s] - if isinstance(env.action_space, gym.spaces.Box): - clipped_actions = torch.clamp(clipped_actions, torch.from_numpy(env.action_space.low), torch.from_numpy(env.action_space.high)) - - o, r, d, _ = env.step(clipped_actions) - states[:, s + 1] = torch.tensor(o).clone().detach() - rewards[:, s] = torch.tensor(r).clone().detach() - dones[:, s + 1] = torch.tensor(d).clone().detach() - - dones = dones.cumsum(1) > 0 - - states = states[:, :max_steps_per_episode] - actions = actions[:, :max_steps_per_episode] - rewards = rewards[:, :max_steps_per_episode] - dones = dones[:, :max_steps_per_episode] - - return states, actions, rewards, dones - - -def rollout_sb3(env, policy, n_episodes, max_steps_per_episode): - states = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.observation_space.shape) - actions = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.action_space.shape) - rewards = torch.zeros(n_episodes, max_steps_per_episode + 1) - dones = torch.ones(n_episodes, max_steps_per_episode + 1, dtype=bool) - - for e in tqdm(range(n_episodes)): - states[e, 0] = torch.tensor(env.reset()).clone().detach() - dones[e, 0] = False - - for s in range(max_steps_per_episode): - action, _ = policy.predict(states[e, s]) - actions[e, s] = torch.tensor(action).clone().detach() - - clipped_actions = actions[e, s] - if isinstance(env.action_space, gym.spaces.Box): - clipped_actions = torch.clamp(clipped_actions, torch.from_numpy(env.action_space.low), torch.from_numpy(env.action_space.high)) - - o, r, d, _ = env.step(clipped_actions) - states[e, s + 1] = torch.tensor(o).clone().detach() - rewards[e, s] = torch.tensor(r).clone().detach() - dones[e, s + 1] = torch.tensor(d).clone().detach() - - if d: - break - - dones = dones.cumsum(1) > 0 - - states = states[:, :max_steps_per_episode] - actions = actions[:, :max_steps_per_episode] - rewards = rewards[:, :max_steps_per_episode] - dones = dones[:, :max_steps_per_episode] - - return states, actions, rewards, dones diff --git a/scratch/etienne/trpo/core/test_optimization.py b/scratch/etienne/trpo/core/test_optimization.py deleted file mode 100644 index aa49bb5..0000000 --- a/scratch/etienne/trpo/core/test_optimization.py +++ /dev/null @@ -1,23 +0,0 @@ -import torch -from optimization import conjugate_gradient - -def test_cg_eye(): - A = torch.eye(2) - b = torch.tensor([1., 2.]) - x1 = conjugate_gradient(lambda x: A @ x, b, 2) - x2 = torch.inverse(A) @ b - assert torch.allclose(x1, x2) - -def test_cg_eyep1(): - A = torch.eye(2) + 1 - b = torch.tensor([1., 2.]) - x1 = conjugate_gradient(lambda x: A @ x, b, 2) - x2 = torch.inverse(A) @ b - assert torch.allclose(x1, x2, atol=1e-7) - -def test_cg3(): - A = torch.tensor([[4., 2.], [2., 4.]]) - b = torch.tensor([2., 1.]) - x1 = conjugate_gradient(lambda x: A @ x, b, 100) - x2 = torch.inverse(A) @ b - assert torch.allclose(x1, x2) diff --git a/scratch/etienne/trpo/core/trpo.py b/scratch/etienne/trpo/core/trpo.py deleted file mode 100644 index d91e363..0000000 --- a/scratch/etienne/trpo/core/trpo.py +++ /dev/null @@ -1,79 +0,0 @@ -import torch -from core.reparam_module import ReparamPolicy -from core.sampling import rollout -from core.value_estimation import gae -from core.optimization import conjugate_gradient, line_search - -def trpo(env_fn, value, policy, epochs, rollout_episodes, rollout_steps, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1): - - policy(torch.zeros(env_fn(0).observation_space.shape)) - policy = ReparamPolicy(policy) - - for epoch in range(epochs): - policy.eval() - states, actions, rewards, dones = rollout(env_fn, policy, rollout_episodes, rollout_steps) - - print('mean', states[~dones].mean(0)) - print('std', states[~dones].std(0)) - - print(f'Iteration {epoch} mean episode length {(~dones).sum() / states.shape[0]}') - print(f'Iteration {epoch} mean reward per episode {rewards[~dones].sum() / states.shape[0]}') - - policy.train() - value.train() - value, policy = trpo_step(value, policy, states, actions, rewards, dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters, cg_damping) - - return value, policy - -def trpo_step(value, policy, states, actions, rewards, dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1): - - states = states.detach() - actions = actions.detach() - rewards = rewards.detach() - dones = dones.detach() - - advantages, returns, valid = gae(states, rewards, value(states), dones, gamma, gae_lambda) - advantages = advantages.detach() - returns = returns.detach() - - # update value function - - for _ in range(v_iters): - v_opt.zero_grad() - value_loss = (value(states) - returns).pow(2)[valid].mean() - value_loss.backward() - v_opt.step() - - # compute policy gradient - - plogprob = policy.log_prob(policy(states), actions) - surrogate_advantage = (plogprob * advantages)[valid].sum() / states.shape[0] - g = torch.cat(torch.autograd.grad(surrogate_advantage, policy.flat_param)).detach() - - def Hx(x): - kl = policy.kl_divergence(policy(states), policy(states).detach())[valid].mean() - dKL = torch.cat(torch.autograd.grad(kl, policy.flat_param, create_graph=True)) - H_x = torch.cat(torch.autograd.grad(dKL.T @ x, policy.flat_param)).detach() - return H_x + cg_damping * x - - x = conjugate_gradient(Hx, g, cg_iters) - npg = torch.sqrt(2 * delta / (x.T @ Hx(x))) * x - - # perform line search - - def L(theta): - rplogprob = policy.log_prob(policy(states, flat_param=theta), actions) - return ((rplogprob - plogprob.detach()).exp() * advantages)[valid].sum() / advantages.shape[0] - - condition = lambda theta: policy.kl_divergence(policy(states, flat_param=theta), policy(states))[valid].mean() < delta - - x0 = policy.flat_param - g0 = torch.cat(torch.autograd.grad(L(x0), x0)) - theta = line_search(L, x0, npg, g0, backtrack_coeff, condition, max_steps=backtrack_iters) - - # update policy parameters - - with torch.no_grad(): - policy.flat_param.copy_(theta) - - return value, policy diff --git a/scratch/etienne/trpo/core/value.py b/scratch/etienne/trpo/core/value.py deleted file mode 100644 index 2be3b78..0000000 --- a/scratch/etienne/trpo/core/value.py +++ /dev/null @@ -1,48 +0,0 @@ -import torch -import torch.nn as nn -from torch.distributions import Normal -from torch.distributions.kl import kl_divergence - -class Value(nn.Module): - - def __init__(self): - super().__init__() - self.nn = nn.Sequential( - nn.LazyLinear(50), - nn.Tanh(), - nn.LazyLinear(50), - nn.Tanh(), - nn.LazyLinear(1), - ) - - def forward(self, states): - return self.nn(states).squeeze(-1) - -class SetValue(Value): - - def forward(self, states): - batch_size = states.shape[:-2] - states = torch.cat((states[..., :1, [0, 1]], states[..., :, [2, 5]]), axis=-2).reshape(*batch_size, -1) - return super().forward(states) - -class DeepSetValue(nn.Module): - - def __init__(self): - super().__init__() - self.elem = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - ) - self.glob = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(1), - ) - - def forward(self, states): - return self.glob(self.elem(states).sum(-2)).squeeze(-1) diff --git a/scratch/etienne/trpo/core/value_estimation.py b/scratch/etienne/trpo/core/value_estimation.py deleted file mode 100644 index b558ac2..0000000 --- a/scratch/etienne/trpo/core/value_estimation.py +++ /dev/null @@ -1,40 +0,0 @@ -from operator import index -import torch - -def gae(states, rewards, values, dones, gamma, gae_lambda): - assert rewards.shape == values.shape == dones.shape - n_episodes, n_steps = rewards.shape - - valid = ~dones - valid[..., -1] = False - - td = rewards + gamma * torch.roll(values, shifts=-1, dims=1) - values - adv = td.repeat(n_steps, 1, 1).transpose(0, 1) - assert adv.shape == (n_episodes, n_steps, n_steps) - - step_start, step = torch.meshgrid(torch.arange(n_steps), torch.arange(n_steps), indexing='ij') - past = step < step_start - - # add up discounted temporal differences - discount = torch.minimum(torch.tensor(gamma).log() * (step - step_start), torch.tensor(0.)).exp() - discount = discount * ~past - discount = discount * valid.unsqueeze(1) - - adv = adv * discount - adv = adv.cumsum(2) # eq. (14) - assert adv.shape == (n_episodes, n_steps, n_steps) - - # add up discounted k-advantages - lambda_discount = torch.minimum(torch.tensor(gae_lambda).log() * (step - step_start), torch.tensor(0.)).exp() - lambda_discount = lambda_discount * ~past - lambda_discount = lambda_discount * valid.unsqueeze(1) - - adv = adv * lambda_discount - adv = adv.sum(2) / (lambda_discount.sum(2) + 1e-10) # eq. (16) - - adv = (adv - adv[valid].mean()) / adv[valid].std() - assert adv.shape == rewards.shape == values.shape - - returns = adv + values - - return adv, returns, valid diff --git a/scratch/etienne/trpo/bc-intersimple-setobs2.py b/scratch/etienne/trpo/experiments/bc-intersimple-setobs2.py similarity index 96% rename from scratch/etienne/trpo/bc-intersimple-setobs2.py rename to scratch/etienne/trpo/experiments/bc-intersimple-setobs2.py index da9c51d..04d2b90 100644 --- a/scratch/etienne/trpo/bc-intersimple-setobs2.py +++ b/scratch/etienne/trpo/experiments/bc-intersimple-setobs2.py @@ -26,7 +26,7 @@ torch.save(policy.state_dict(), 'bc-intersimple-setobs2.pt') # %% import numpy as np from core.policy import SetPolicy -from wrappers import Setobs, TransformObservation, CollisionPenaltyWrapper +from util.wrappers import Setobs, TransformObservation, CollisionPenaltyWrapper from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools diff --git a/scratch/etienne/trpo/gail-intersimple-minobs.py b/scratch/etienne/trpo/experiments/gail-intersimple-minobs.py similarity index 97% rename from scratch/etienne/trpo/gail-intersimple-minobs.py rename to scratch/etienne/trpo/experiments/gail-intersimple-minobs.py index bfeed96..62d4a2e 100644 --- a/scratch/etienne/trpo/gail-intersimple-minobs.py +++ b/scratch/etienne/trpo/experiments/gail-intersimple-minobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Minobs +from util.wrappers import CollisionPenaltyWrapper, Minobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/gail-intersimple-minobs2.py b/scratch/etienne/trpo/experiments/gail-intersimple-minobs2.py similarity index 97% rename from scratch/etienne/trpo/gail-intersimple-minobs2.py rename to scratch/etienne/trpo/experiments/gail-intersimple-minobs2.py index 0e1cf76..008b8ca 100644 --- a/scratch/etienne/trpo/gail-intersimple-minobs2.py +++ b/scratch/etienne/trpo/experiments/gail-intersimple-minobs2.py @@ -8,7 +8,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Minobs +from util.wrappers import CollisionPenaltyWrapper, Minobs import numpy as np from gym.wrappers import TransformObservation from core.reparam_module import ReparamPolicy diff --git a/scratch/etienne/trpo/gail-intersimple-normobs.py b/scratch/etienne/trpo/experiments/gail-intersimple-normobs.py similarity index 97% rename from scratch/etienne/trpo/gail-intersimple-normobs.py rename to scratch/etienne/trpo/experiments/gail-intersimple-normobs.py index 1ed0bd4..881efa0 100644 --- a/scratch/etienne/trpo/gail-intersimple-normobs.py +++ b/scratch/etienne/trpo/experiments/gail-intersimple-normobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/gail-intersimple-setobs.py b/scratch/etienne/trpo/experiments/gail-intersimple-setobs.py similarity index 97% rename from scratch/etienne/trpo/gail-intersimple-setobs.py rename to scratch/etienne/trpo/experiments/gail-intersimple-setobs.py index add36e1..7be8299 100644 --- a/scratch/etienne/trpo/gail-intersimple-setobs.py +++ b/scratch/etienne/trpo/experiments/gail-intersimple-setobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Setobs +from util.wrappers import CollisionPenaltyWrapper, Setobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/gail-intersimple-setobs2-recurrent.py b/scratch/etienne/trpo/experiments/gail-intersimple-setobs2-recurrent.py similarity index 97% rename from scratch/etienne/trpo/gail-intersimple-setobs2-recurrent.py rename to scratch/etienne/trpo/experiments/gail-intersimple-setobs2-recurrent.py index 69732e6..2130b6f 100644 --- a/scratch/etienne/trpo/gail-intersimple-setobs2-recurrent.py +++ b/scratch/etienne/trpo/experiments/gail-intersimple-setobs2-recurrent.py @@ -8,7 +8,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Setobs +from util.wrappers import CollisionPenaltyWrapper, Setobs import numpy as np from gym.wrappers import TransformObservation from core.reparam_module import ReparamPolicy diff --git a/scratch/etienne/trpo/gail-intersimple-setobs2.py b/scratch/etienne/trpo/experiments/gail-intersimple-setobs2.py similarity index 98% rename from scratch/etienne/trpo/gail-intersimple-setobs2.py rename to scratch/etienne/trpo/experiments/gail-intersimple-setobs2.py index 5a0c51b..e122ee6 100644 --- a/scratch/etienne/trpo/gail-intersimple-setobs2.py +++ b/scratch/etienne/trpo/experiments/gail-intersimple-setobs2.py @@ -8,7 +8,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Setobs +from util.wrappers import CollisionPenaltyWrapper, Setobs import numpy as np from gym.wrappers import TransformObservation from core.reparam_module import ReparamPolicy diff --git a/scratch/etienne/trpo/gail-intersimple.py b/scratch/etienne/trpo/experiments/gail-intersimple.py similarity index 96% rename from scratch/etienne/trpo/gail-intersimple.py rename to scratch/etienne/trpo/experiments/gail-intersimple.py index 4fac87e..11b050d 100644 --- a/scratch/etienne/trpo/gail-intersimple.py +++ b/scratch/etienne/trpo/experiments/gail-intersimple.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat( n_rays=5, diff --git a/scratch/etienne/trpo/gail-options-minobs.py b/scratch/etienne/trpo/experiments/gail-options-minobs.py similarity index 97% rename from scratch/etienne/trpo/gail-options-minobs.py rename to scratch/etienne/trpo/experiments/gail-options-minobs.py index ca9ef13..b738261 100644 --- a/scratch/etienne/trpo/gail-options-minobs.py +++ b/scratch/etienne/trpo/experiments/gail-options-minobs.py @@ -8,7 +8,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter diff --git a/scratch/etienne/trpo/gail-options-setobs.py b/scratch/etienne/trpo/experiments/gail-options-setobs.py similarity index 97% rename from scratch/etienne/trpo/gail-options-setobs.py rename to scratch/etienne/trpo/experiments/gail-options-setobs.py index c9ca49c..288647b 100644 --- a/scratch/etienne/trpo/gail-options-setobs.py +++ b/scratch/etienne/trpo/experiments/gail-options-setobs.py @@ -8,7 +8,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter diff --git a/scratch/etienne/trpo/gail-options-setobs2.py b/scratch/etienne/trpo/experiments/gail-options-setobs2.py similarity index 89% rename from scratch/etienne/trpo/gail-options-setobs2.py rename to scratch/etienne/trpo/experiments/gail-options-setobs2.py index 267b179..953b347 100644 --- a/scratch/etienne/trpo/gail-options-setobs2.py +++ b/scratch/etienne/trpo/experiments/gail-options-setobs2.py @@ -9,7 +9,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter @@ -56,6 +56,11 @@ expert_data = torch.load('intersimple-expert-data-setobs2.pt') expert_data = Buffer(*expert_data) # %% +def callback(epoch, value, policy): + if not epoch % 10: + torch.save(policy.state_dict(), f'gail-options-setobs2-{epoch}.pt') + torch.save(value.state_dict(), f'gail-options-setobs2-value-{epoch}.pt') + value, policy = gail( env_fn=env_fn, expert_data=expert_data, @@ -66,7 +71,7 @@ value, policy = gail( value=value, v_opt=v_opt, v_iters=1000, - epochs=200, + epochs=300, rollout_episodes=60, rollout_steps=60, gamma=0.99, @@ -75,6 +80,7 @@ value, policy = gail( backtrack_coeff=0.8, backtrack_iters=10, logger=SummaryWriter(comment='gail-options-setobs2'), + callback=callback, ) torch.save(policy.state_dict(), 'gail-options-setobs2.pt') diff --git a/scratch/etienne/trpo/gail-pendulum.py b/scratch/etienne/trpo/experiments/gail-pendulum.py similarity index 100% rename from scratch/etienne/trpo/gail-pendulum.py rename to scratch/etienne/trpo/experiments/gail-pendulum.py diff --git a/scratch/etienne/trpo/gail-ppo-intersimple-minobs.py b/scratch/etienne/trpo/experiments/gail-ppo-intersimple-minobs.py similarity index 97% rename from scratch/etienne/trpo/gail-ppo-intersimple-minobs.py rename to scratch/etienne/trpo/experiments/gail-ppo-intersimple-minobs.py index 46bcb18..21afb69 100644 --- a/scratch/etienne/trpo/gail-ppo-intersimple-minobs.py +++ b/scratch/etienne/trpo/experiments/gail-ppo-intersimple-minobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Minobs +from util.wrappers import CollisionPenaltyWrapper, Minobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/gail-ppo-intersimple-normobs.py b/scratch/etienne/trpo/experiments/gail-ppo-intersimple-normobs.py similarity index 97% rename from scratch/etienne/trpo/gail-ppo-intersimple-normobs.py rename to scratch/etienne/trpo/experiments/gail-ppo-intersimple-normobs.py index 028328a..329fb46 100644 --- a/scratch/etienne/trpo/gail-ppo-intersimple-normobs.py +++ b/scratch/etienne/trpo/experiments/gail-ppo-intersimple-normobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/gail-ppo-intersimple-setobs2.py b/scratch/etienne/trpo/experiments/gail-ppo-intersimple-setobs2.py similarity index 98% rename from scratch/etienne/trpo/gail-ppo-intersimple-setobs2.py rename to scratch/etienne/trpo/experiments/gail-ppo-intersimple-setobs2.py index 4715256..1de435d 100644 --- a/scratch/etienne/trpo/gail-ppo-intersimple-setobs2.py +++ b/scratch/etienne/trpo/experiments/gail-ppo-intersimple-setobs2.py @@ -8,7 +8,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Setobs +from util.wrappers import CollisionPenaltyWrapper, Setobs import numpy as np from gym.wrappers import TransformObservation from core.reparam_module import ReparamPolicy diff --git a/scratch/etienne/trpo/gail-ppo-intersimple.py b/scratch/etienne/trpo/experiments/gail-ppo-intersimple.py similarity index 96% rename from scratch/etienne/trpo/gail-ppo-intersimple.py rename to scratch/etienne/trpo/experiments/gail-ppo-intersimple.py index 833c028..7412b3f 100644 --- a/scratch/etienne/trpo/gail-ppo-intersimple.py +++ b/scratch/etienne/trpo/experiments/gail-ppo-intersimple.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat( n_rays=5, diff --git a/scratch/etienne/trpo/gail-ppo-options-minobs.py b/scratch/etienne/trpo/experiments/gail-ppo-options-minobs.py similarity index 97% rename from scratch/etienne/trpo/gail-ppo-options-minobs.py rename to scratch/etienne/trpo/experiments/gail-ppo-options-minobs.py index b264135..f25a9ea 100644 --- a/scratch/etienne/trpo/gail-ppo-options-minobs.py +++ b/scratch/etienne/trpo/experiments/gail-ppo-options-minobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter diff --git a/scratch/etienne/trpo/gail-ppo-options-setobs.py b/scratch/etienne/trpo/experiments/gail-ppo-options-setobs.py similarity index 97% rename from scratch/etienne/trpo/gail-ppo-options-setobs.py rename to scratch/etienne/trpo/experiments/gail-ppo-options-setobs.py index 4f2d463..8fa9339 100644 --- a/scratch/etienne/trpo/gail-ppo-options-setobs.py +++ b/scratch/etienne/trpo/experiments/gail-ppo-options-setobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter diff --git a/scratch/etienne/trpo/gail-ppo-options-setobs2.py b/scratch/etienne/trpo/experiments/gail-ppo-options-setobs2.py similarity index 89% rename from scratch/etienne/trpo/gail-ppo-options-setobs2.py rename to scratch/etienne/trpo/experiments/gail-ppo-options-setobs2.py index 7d19fc1..2ad3a46 100644 --- a/scratch/etienne/trpo/gail-ppo-options-setobs2.py +++ b/scratch/etienne/trpo/experiments/gail-ppo-options-setobs2.py @@ -8,7 +8,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter @@ -57,6 +57,11 @@ expert_data = torch.load('intersimple-expert-data-setobs2.pt') expert_data = Buffer(*expert_data) # %% +def callback(epoch, value, policy): + if not epoch % 10: + torch.save(policy.state_dict(), f'gail-ppo-options-setobs2-{epoch}.pt') + torch.save(value.state_dict(), f'gail-ppo-options-setobs2-value-{epoch}.pt') + value, policy = gail_ppo( env_fn=env_fn, expert_data=expert_data, @@ -76,6 +81,7 @@ value, policy = gail_ppo( pi_opt=pi_opt, pi_iters=100, logger=SummaryWriter(comment='gail-ppo-options-setobs2'), + callback=callback, ) torch.save(policy.state_dict(), 'gail-ppo-options-setobs2.pt') diff --git a/scratch/etienne/trpo/intersimple-expert-action-profiles.ipynb b/scratch/etienne/trpo/experiments/intersimple-expert-action-profiles.ipynb similarity index 100% rename from scratch/etienne/trpo/intersimple-expert-action-profiles.ipynb rename to scratch/etienne/trpo/experiments/intersimple-expert-action-profiles.ipynb diff --git a/scratch/etienne/trpo/intersimple-expert-rollout-minobs.py b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-minobs.py similarity index 96% rename from scratch/etienne/trpo/intersimple-expert-rollout-minobs.py rename to scratch/etienne/trpo/experiments/intersimple-expert-rollout-minobs.py index 0a00f67..9cb3802 100644 --- a/scratch/etienne/trpo/intersimple-expert-rollout-minobs.py +++ b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-minobs.py @@ -4,7 +4,7 @@ from core.sampling import rollout_sb3 from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward from intersim.expert import NormalizedIntersimpleExpert -from wrappers import CollisionPenaltyWrapper, Minobs +from util.wrappers import CollisionPenaltyWrapper, Minobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/intersimple-expert-rollout-minobs2.py b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-minobs2.py similarity index 96% rename from scratch/etienne/trpo/intersimple-expert-rollout-minobs2.py rename to scratch/etienne/trpo/experiments/intersimple-expert-rollout-minobs2.py index b65ad11..e839f74 100644 --- a/scratch/etienne/trpo/intersimple-expert-rollout-minobs2.py +++ b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-minobs2.py @@ -4,7 +4,7 @@ from core.sampling import rollout_sb3 from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward from intersim.expert import NormalizedIntersimpleExpert -from wrappers import CollisionPenaltyWrapper, Minobs +from util.wrappers import CollisionPenaltyWrapper, Minobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/intersimple-expert-rollout-normobs.py b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-normobs.py similarity index 97% rename from scratch/etienne/trpo/intersimple-expert-rollout-normobs.py rename to scratch/etienne/trpo/experiments/intersimple-expert-rollout-normobs.py index e9340ae..b839a3a 100644 --- a/scratch/etienne/trpo/intersimple-expert-rollout-normobs.py +++ b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-normobs.py @@ -4,7 +4,7 @@ from core.sampling import rollout_sb3 from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward from intersim.expert import NormalizedIntersimpleExpert -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/intersimple-expert-rollout-setobs.py b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-setobs.py similarity index 96% rename from scratch/etienne/trpo/intersimple-expert-rollout-setobs.py rename to scratch/etienne/trpo/experiments/intersimple-expert-rollout-setobs.py index 0ab531b..dcf5223 100644 --- a/scratch/etienne/trpo/intersimple-expert-rollout-setobs.py +++ b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-setobs.py @@ -4,7 +4,7 @@ from core.sampling import rollout_sb3 from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward from intersim.expert import NormalizedIntersimpleExpert -from wrappers import CollisionPenaltyWrapper, Setobs +from util.wrappers import CollisionPenaltyWrapper, Setobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/intersimple-expert-rollout-setobs2.py b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-setobs2.py similarity index 96% rename from scratch/etienne/trpo/intersimple-expert-rollout-setobs2.py rename to scratch/etienne/trpo/experiments/intersimple-expert-rollout-setobs2.py index 16ecd67..28139f2 100644 --- a/scratch/etienne/trpo/intersimple-expert-rollout-setobs2.py +++ b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-setobs2.py @@ -4,7 +4,7 @@ from core.sampling import rollout_sb3 from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward from intersim.expert import NormalizedIntersimpleExpert -from wrappers import CollisionPenaltyWrapper, Setobs +from util.wrappers import CollisionPenaltyWrapper, Setobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/intersimple-expert-rollout.py b/scratch/etienne/trpo/experiments/intersimple-expert-rollout.py similarity index 94% rename from scratch/etienne/trpo/intersimple-expert-rollout.py rename to scratch/etienne/trpo/experiments/intersimple-expert-rollout.py index ac0501c..d3a7deb 100644 --- a/scratch/etienne/trpo/intersimple-expert-rollout.py +++ b/scratch/etienne/trpo/experiments/intersimple-expert-rollout.py @@ -4,7 +4,7 @@ from core.sampling import rollout_sb3 from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward from intersim.expert import NormalizedIntersimpleExpert -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper env = CollisionPenaltyWrapper(IntersimpleLidarFlat( n_rays=5, diff --git a/scratch/etienne/trpo/ppo-intersimple-minobs.py b/scratch/etienne/trpo/experiments/ppo-intersimple-minobs.py similarity index 98% rename from scratch/etienne/trpo/ppo-intersimple-minobs.py rename to scratch/etienne/trpo/experiments/ppo-intersimple-minobs.py index 89f61a6..3c648ce 100644 --- a/scratch/etienne/trpo/ppo-intersimple-minobs.py +++ b/scratch/etienne/trpo/experiments/ppo-intersimple-minobs.py @@ -9,7 +9,7 @@ import torch.optim import numpy as np from gym.wrappers import TransformObservation -from wrappers import Minobs +from util.wrappers import Minobs obs_min = np.array([ [-1000, -1000, 0, -np.pi, -1e-1, 0.], diff --git a/scratch/etienne/trpo/ppo-intersimple-minobs2.py b/scratch/etienne/trpo/experiments/ppo-intersimple-minobs2.py similarity index 98% rename from scratch/etienne/trpo/ppo-intersimple-minobs2.py rename to scratch/etienne/trpo/experiments/ppo-intersimple-minobs2.py index 5bed6b3..f119e74 100644 --- a/scratch/etienne/trpo/ppo-intersimple-minobs2.py +++ b/scratch/etienne/trpo/experiments/ppo-intersimple-minobs2.py @@ -9,7 +9,7 @@ import torch.optim import numpy as np from gym.wrappers import TransformObservation -from wrappers import Minobs +from util.wrappers import Minobs obs_min = np.array([ [-1000, -1000, 0, -np.pi, -1e-1, 0.], diff --git a/scratch/etienne/trpo/ppo-intersimple-normobs.py b/scratch/etienne/trpo/experiments/ppo-intersimple-normobs.py similarity index 100% rename from scratch/etienne/trpo/ppo-intersimple-normobs.py rename to scratch/etienne/trpo/experiments/ppo-intersimple-normobs.py diff --git a/scratch/etienne/trpo/ppo-intersimple.py b/scratch/etienne/trpo/experiments/ppo-intersimple.py similarity index 100% rename from scratch/etienne/trpo/ppo-intersimple.py rename to scratch/etienne/trpo/experiments/ppo-intersimple.py diff --git a/scratch/etienne/trpo/ppo-options-minobs.py b/scratch/etienne/trpo/experiments/ppo-options-minobs.py similarity index 95% rename from scratch/etienne/trpo/ppo-options-minobs.py rename to scratch/etienne/trpo/experiments/ppo-options-minobs.py index 38aa899..009f41d 100644 --- a/scratch/etienne/trpo/ppo-options-minobs.py +++ b/scratch/etienne/trpo/experiments/ppo-options-minobs.py @@ -9,9 +9,9 @@ from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools import numpy as np -from wrappers import CollisionPenaltyWrapper, TransformObservation +from util.wrappers import CollisionPenaltyWrapper, TransformObservation -from wrappers import Minobs +from util.wrappers import Minobs from options.options import OptionsEnv obs_min = np.array([ diff --git a/scratch/etienne/trpo/ppo-pendulum.py b/scratch/etienne/trpo/experiments/ppo-pendulum.py similarity index 100% rename from scratch/etienne/trpo/ppo-pendulum.py rename to scratch/etienne/trpo/experiments/ppo-pendulum.py diff --git a/scratch/etienne/trpo/readme.md b/scratch/etienne/trpo/experiments/readme.md similarity index 100% rename from scratch/etienne/trpo/readme.md rename to scratch/etienne/trpo/experiments/readme.md diff --git a/scratch/etienne/trpo/experiments/requirements.txt b/scratch/etienne/trpo/experiments/requirements.txt new file mode 100644 index 0000000..bd1ffb4 --- /dev/null +++ b/scratch/etienne/trpo/experiments/requirements.txt @@ -0,0 +1,3 @@ +torch +stable-baselines3 +gym diff --git a/scratch/etienne/trpo/experiments/sgail-options-setobs2.py b/scratch/etienne/trpo/experiments/sgail-options-setobs2.py new file mode 100644 index 0000000..02c8bff --- /dev/null +++ b/scratch/etienne/trpo/experiments/sgail-options-setobs2.py @@ -0,0 +1,108 @@ +# %% +import gym +from safe_options.options import gail +from core.gail import Buffer +from core.value import SetValue +from safe_options.policy import SetMaskedDiscretePolicy +from core.discriminator import DeepsetDiscriminator +import torch.optim +from intersim.envs import IntersimpleLidarFlatRandom +from intersim.envs.intersimple import speed_reward +import functools +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +import numpy as np +from safe_options.options import SafeOptionsEnv +from torch.utils.tensorboard import SummaryWriter +from core.reparam_module import ReparamPolicy + +obs_min = np.array([ + [-1000, -1000, 0, -np.pi, -1e-1, 0.], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], +]).reshape(-1) + +obs_max = np.array([ + [1000, 1000, 20, np.pi, 1e-1, 0.], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], +]).reshape(-1) + +envs = [SafeOptionsEnv(Setobs( + TransformObservation(CollisionPenaltyWrapper(IntersimpleLidarFlatRandom( + n_rays=5, + reward=functools.partial( + speed_reward, + collision_penalty=0 + ), + stop_on_collision=True, + ), collision_distance=6, collision_penalty=100), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) +), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5), (10, 5)], safe_actions_collision_method='circle', abort_unsafe_collision_method='circle') for _ in range(60)] + +env_fn = lambda i: envs[i] +policy = SetMaskedDiscretePolicy(env_fn(0).action_space.n) +value = SetValue() +v_opt = torch.optim.Adam(value.parameters(), lr=1e-4) + +discriminator = DeepsetDiscriminator() +disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-4) + +expert_data = torch.load('intersimple-expert-data-setobs2.pt') +expert_data = Buffer(*expert_data) + +# %% +def callback(epoch, value, policy): + if not epoch % 10: + torch.save(policy.state_dict(), f'sgail-options-setobs2-{epoch}.pt') + torch.save(value.state_dict(), f'sgail-options-setobs2-value-{epoch}.pt') + +value, policy = gail( + env_fn=env_fn, + expert_data=expert_data, + discriminator=discriminator, + disc_opt=disc_opt, + disc_iters=100, + policy=policy, + value=value, + v_opt=v_opt, + v_iters=1000, + epochs=300, + rollout_episodes=60, + rollout_steps=60, + gamma=0.99, + gae_lambda=0.9, + delta=0.01, + backtrack_coeff=0.8, + backtrack_iters=10, + logger=SummaryWriter(comment='sgail-options-setobs2'), + callback=callback, +) + +torch.save(policy.state_dict(), 'sgail-options-setobs2.pt') + +# %% +policy = SetMaskedDiscretePolicy(env_fn(0).action_space.n) +policy(torch.zeros(env_fn(0).observation_space['observation'].shape), torch.zeros(env_fn(0).observation_space['safe_actions'].shape)) +policy = ReparamPolicy(policy) +policy.load_state_dict(torch.load('sgail-options-setobs2.pt')) + +env = env_fn(0) +obs = env.reset() +env.render(mode='post') +for i in range(300): + action = policy.sample(policy( + torch.tensor(obs['observation'], dtype=torch.float32), + torch.tensor(obs['safe_actions'], dtype=torch.float32), + )) + obs, reward, done, _ = env.step(action, render_mode='post') + print('step', i, 'reward', reward) + if done: + break +env.close() + +# %% diff --git a/scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py b/scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py new file mode 100644 index 0000000..5365a94 --- /dev/null +++ b/scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py @@ -0,0 +1,107 @@ +# %% +import gym +from safe_options.options import gail_ppo, Buffer +from core.value import SetValue +from safe_options.policy import SetMaskedDiscretePolicy +from core.discriminator import DeepsetDiscriminator +import torch.optim +from intersim.envs import IntersimpleLidarFlatRandom +from intersim.envs.intersimple import speed_reward +import functools +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +import numpy as np +from safe_options.options import SafeOptionsEnv +from torch.utils.tensorboard import SummaryWriter + +obs_min = np.array([ + [-1000, -1000, 0, -np.pi, -1e-1, 0.], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], +]).reshape(-1) + +obs_max = np.array([ + [1000, 1000, 20, np.pi, 1e-1, 0.], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], +]).reshape(-1) + +envs = [SafeOptionsEnv(Setobs( + TransformObservation(CollisionPenaltyWrapper(IntersimpleLidarFlatRandom( + n_rays=5, + reward=functools.partial( + speed_reward, + collision_penalty=0 + ), + stop_on_collision=True, + ), collision_distance=6, collision_penalty=100), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) +), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5), (10, 5)], safe_actions_collision_method='circle', abort_unsafe_collision_method='circle') for _ in range(60)] + +env_fn = lambda i: envs[i] + +policy = SetMaskedDiscretePolicy(env_fn(0).action_space.n) +pi_opt = torch.optim.Adam(policy.parameters(), lr=3e-4) + +value = SetValue() +v_opt = torch.optim.Adam(value.parameters(), lr=1e-3) + +discriminator = DeepsetDiscriminator() +disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-4) + +expert_data = torch.load('intersimple-expert-data-setobs2.pt') +expert_data = Buffer(*expert_data) + +# %% +def callback(epoch, value, policy): + if not epoch % 10: + torch.save(policy.state_dict(), f'sgail-ppo-options-setobs2-{epoch}.pt') + torch.save(value.state_dict(), f'sgail-ppo-options-setobs2-value-{epoch}.pt') + +value, policy = gail_ppo( + env_fn=env_fn, + expert_data=expert_data, + discriminator=discriminator, + disc_opt=disc_opt, + disc_iters=100, + policy=policy, + value=value, + v_opt=v_opt, + v_iters=1000, + epochs=200, + rollout_episodes=60, + rollout_steps=60, + gamma=0.99, + gae_lambda=0.9, + clip_ratio=0.2, + pi_opt=pi_opt, + pi_iters=100, + logger=SummaryWriter(comment='sgail-ppo-options-setobs2'), + callback=callback, +) + +torch.save(policy.state_dict(), 'sgail-ppo-options-setobs2.pt') + +# %% +policy = SetMaskedDiscretePolicy(env_fn(0).action_space.n) +policy(torch.zeros(env_fn(0).observation_space['observation'].shape), torch.zeros(env_fn(0).observation_space['safe_actions'].shape)) +policy.load_state_dict(torch.load('sgail-ppo-options-setobs2.pt')) + +env = env_fn(0) +obs = env.reset() +env.render(mode='post') +for i in range(300): + action = policy.sample(policy( + torch.tensor(obs['observation'], dtype=torch.float32), + torch.tensor(obs['safe_actions'], dtype=torch.float32), + )) + obs, reward, done, _ = env.step(action, render_mode='post') + print('step', i, 'reward', reward, 'safe actions', obs['safe_actions']) + if done: + break +env.close() +# %% diff --git a/scratch/etienne/trpo/trpo-intersimple-minobs.py b/scratch/etienne/trpo/experiments/trpo-intersimple-minobs.py similarity index 96% rename from scratch/etienne/trpo/trpo-intersimple-minobs.py rename to scratch/etienne/trpo/experiments/trpo-intersimple-minobs.py index 93081d9..b3dc363 100644 --- a/scratch/etienne/trpo/trpo-intersimple-minobs.py +++ b/scratch/etienne/trpo/experiments/trpo-intersimple-minobs.py @@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward import functools import numpy as np from gym.wrappers import TransformObservation -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper from core.reparam_module import ReparamPolicy -from wrappers import Minobs +from util.wrappers import Minobs obs_min = np.array([ [-1000, -1000, 0, -np.pi, -1e-1, 0.], diff --git a/scratch/etienne/trpo/trpo-intersimple-minobs2.py b/scratch/etienne/trpo/experiments/trpo-intersimple-minobs2.py similarity index 96% rename from scratch/etienne/trpo/trpo-intersimple-minobs2.py rename to scratch/etienne/trpo/experiments/trpo-intersimple-minobs2.py index 58f6b28..0321938 100644 --- a/scratch/etienne/trpo/trpo-intersimple-minobs2.py +++ b/scratch/etienne/trpo/experiments/trpo-intersimple-minobs2.py @@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward import functools import numpy as np from gym.wrappers import TransformObservation -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper from core.reparam_module import ReparamPolicy -from wrappers import Minobs +from util.wrappers import Minobs obs_min = np.array([ [-1000, -1000, 0, -np.pi, -1e-1, 0.], diff --git a/scratch/etienne/trpo/trpo-intersimple-normobs.py b/scratch/etienne/trpo/experiments/trpo-intersimple-normobs.py similarity index 100% rename from scratch/etienne/trpo/trpo-intersimple-normobs.py rename to scratch/etienne/trpo/experiments/trpo-intersimple-normobs.py diff --git a/scratch/etienne/trpo/trpo-intersimple-setobs.py b/scratch/etienne/trpo/experiments/trpo-intersimple-setobs.py similarity index 96% rename from scratch/etienne/trpo/trpo-intersimple-setobs.py rename to scratch/etienne/trpo/experiments/trpo-intersimple-setobs.py index 62545ee..0dd77d9 100644 --- a/scratch/etienne/trpo/trpo-intersimple-setobs.py +++ b/scratch/etienne/trpo/experiments/trpo-intersimple-setobs.py @@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward import functools import numpy as np from gym.wrappers import TransformObservation -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper from core.reparam_module import ReparamPolicy -from wrappers import Setobs +from util.wrappers import Setobs obs_min = np.array([ [-1000, -1000, 0, -np.pi, -1e-1, 0.], diff --git a/scratch/etienne/trpo/trpo-intersimple-setobs2.py b/scratch/etienne/trpo/experiments/trpo-intersimple-setobs2.py similarity index 96% rename from scratch/etienne/trpo/trpo-intersimple-setobs2.py rename to scratch/etienne/trpo/experiments/trpo-intersimple-setobs2.py index 778e9d3..a64e410 100644 --- a/scratch/etienne/trpo/trpo-intersimple-setobs2.py +++ b/scratch/etienne/trpo/experiments/trpo-intersimple-setobs2.py @@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward import functools import numpy as np from gym.wrappers import TransformObservation -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper from core.reparam_module import ReparamPolicy -from wrappers import Setobs +from util.wrappers import Setobs obs_min = np.array([ [-1000, -1000, 0, -np.pi, -1e-1, 0.], diff --git a/scratch/etienne/trpo/trpo-intersimple.py b/scratch/etienne/trpo/experiments/trpo-intersimple.py similarity index 100% rename from scratch/etienne/trpo/trpo-intersimple.py rename to scratch/etienne/trpo/experiments/trpo-intersimple.py diff --git a/scratch/etienne/trpo/trpo-options-minobs.py b/scratch/etienne/trpo/experiments/trpo-options-minobs.py similarity index 95% rename from scratch/etienne/trpo/trpo-options-minobs.py rename to scratch/etienne/trpo/experiments/trpo-options-minobs.py index eaee086..dbc5a09 100644 --- a/scratch/etienne/trpo/trpo-options-minobs.py +++ b/scratch/etienne/trpo/experiments/trpo-options-minobs.py @@ -9,10 +9,10 @@ from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools import numpy as np -from wrappers import CollisionPenaltyWrapper, TransformObservation +from util.wrappers import CollisionPenaltyWrapper, TransformObservation from core.reparam_module import ReparamPolicy -from wrappers import Minobs +from util.wrappers import Minobs from options.options import OptionsEnv obs_min = np.array([ diff --git a/scratch/etienne/trpo/trpo-pendulum-rollout.py b/scratch/etienne/trpo/experiments/trpo-pendulum-rollout.py similarity index 100% rename from scratch/etienne/trpo/trpo-pendulum-rollout.py rename to scratch/etienne/trpo/experiments/trpo-pendulum-rollout.py diff --git a/scratch/etienne/trpo/trpo-pendulum.py b/scratch/etienne/trpo/experiments/trpo-pendulum.py similarity index 100% rename from scratch/etienne/trpo/trpo-pendulum.py rename to scratch/etienne/trpo/experiments/trpo-pendulum.py diff --git a/scratch/etienne/trpo/trpo-walker.py b/scratch/etienne/trpo/experiments/trpo-walker.py similarity index 100% rename from scratch/etienne/trpo/trpo-walker.py rename to scratch/etienne/trpo/experiments/trpo-walker.py diff --git a/scratch/etienne/trpo/experiments/vec-env.ipynb b/scratch/etienne/trpo/experiments/vec-env.ipynb new file mode 100644 index 0000000..1569b69 --- /dev/null +++ b/scratch/etienne/trpo/experiments/vec-env.ipynb @@ -0,0 +1,346 @@ +{ + "cells": [ + { + "cell_type": "code", + "execution_count": 3, + "metadata": {}, + "outputs": [], + "source": [ + "from stable_baselines3.common.env_util import make_vec_env\n", + "import numpy as np" + ] + }, + { + "cell_type": "code", + "execution_count": 4, + "metadata": {}, + "outputs": [], + "source": [ + "env = make_vec_env('Pendulum-v0', n_envs=6)" + ] + }, + { + "cell_type": "code", + "execution_count": 5, + "metadata": {}, + "outputs": [ + { + "data": { + "text/plain": [ + "(6, 3)" + ] + }, + "execution_count": 5, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "obs = env.reset()\n", + "obs.shape" + ] + }, + { + "cell_type": "code", + "execution_count": 6, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "1\n", + "2\n", + "3\n", + "4\n", + "5\n", + "6\n", + "7\n", + "8\n", + "9\n", + "10\n", + "11\n", + "12\n", + "13\n", + "14\n", + "15\n", + "16\n", + "17\n", + "18\n", + "19\n", + "20\n", + "21\n", + "22\n", + "23\n", + "24\n", + "25\n", + "26\n", + "27\n", + "28\n", + "29\n", + "30\n", + "31\n", + "32\n", + "33\n", + "34\n", + "35\n", + "36\n", + "37\n", + "38\n", + "39\n", + "40\n", + "41\n", + "42\n", + "43\n", + "44\n", + "45\n", + "46\n", + "47\n", + "48\n", + "49\n", + "50\n", + "51\n", + "52\n", + "53\n", + "54\n", + "55\n", + "56\n", + "57\n", + "58\n", + "59\n", + "60\n", + "61\n", + "62\n", + "63\n", + "64\n", + "65\n", + "66\n", + "67\n", + "68\n", + "69\n", + "70\n", + "71\n", + "72\n", + "73\n", + "74\n", + "75\n", + "76\n", + "77\n", + "78\n", + "79\n", + "80\n", + "81\n", + "82\n", + "83\n", + "84\n", + "85\n", + "86\n", + "87\n", + "88\n", + "89\n", + "90\n", + "91\n", + "92\n", + "93\n", + "94\n", + "95\n", + "96\n", + "97\n", + "98\n", + "99\n", + "100\n", + "101\n", + "102\n", + "103\n", + "104\n", + "105\n", + "106\n", + "107\n", + "108\n", + "109\n", + "110\n", + "111\n", + "112\n", + "113\n", + "114\n", + "115\n", + "116\n", + "117\n", + "118\n", + "119\n", + "120\n", + "121\n", + "122\n", + "123\n", + "124\n", + "125\n", + "126\n", + "127\n", + "128\n", + "129\n", + "130\n", + "131\n", + "132\n", + "133\n", + "134\n", + "135\n", + "136\n", + "137\n", + "138\n", + "139\n", + "140\n", + "141\n", + "142\n", + "143\n", + "144\n", + "145\n", + "146\n", + "147\n", + "148\n", + "149\n", + "150\n", + "151\n", + "152\n", + "153\n", + "154\n", + "155\n", + "156\n", + "157\n", + "158\n", + "159\n", + "160\n", + "161\n", + "162\n", + "163\n", + "164\n", + "165\n", + "166\n", + "167\n", + "168\n", + "169\n", + "170\n", + "171\n", + "172\n", + "173\n", + "174\n", + "175\n", + "176\n", + "177\n", + "178\n", + "179\n", + "180\n", + "181\n", + "182\n", + "183\n", + "184\n", + "185\n", + "186\n", + "187\n", + "188\n", + "189\n", + "190\n", + "191\n", + "192\n", + "193\n", + "194\n", + "195\n", + "196\n", + "197\n", + "198\n", + "199\n", + "200\n" + ] + } + ], + "source": [ + "dones = [False]\n", + "i = 0\n", + "while not any(dones):\n", + " i += 1\n", + " print(i)\n", + " _, _, dones, _ = env.step(np.zeros((6, 1)))" + ] + }, + { + "cell_type": "code", + "execution_count": 7, + "metadata": {}, + "outputs": [ + { + "data": { + "text/plain": [ + "array([ True, True, True, True, True, True])" + ] + }, + "execution_count": 7, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "dones" + ] + }, + { + "cell_type": "code", + "execution_count": 8, + "metadata": {}, + "outputs": [], + "source": [ + "_, _, dones, _ = env.step(np.zeros((6, 1)))" + ] + }, + { + "cell_type": "code", + "execution_count": 9, + "metadata": {}, + "outputs": [ + { + "data": { + "text/plain": [ + "array([False, False, False, False, False, False])" + ] + }, + "execution_count": 9, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "dones" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [] + } + ], + "metadata": { + "interpreter": { + "hash": "6c7a4ac80dd345f83235e10baa3acc437d966916e1cc075a45b91bb9cc030938" + }, + "kernelspec": { + "display_name": "Python 3.9.7 64-bit ('.venv': venv)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.9.7" + }, + "orig_nbformat": 4 + }, + "nbformat": 4, + "nbformat_minor": 2 +} diff --git a/scratch/etienne/trpo/wgail-intersimple-minobs.py b/scratch/etienne/trpo/experiments/wgail-intersimple-minobs.py similarity index 97% rename from scratch/etienne/trpo/wgail-intersimple-minobs.py rename to scratch/etienne/trpo/experiments/wgail-intersimple-minobs.py index b36988e..677cfef 100644 --- a/scratch/etienne/trpo/wgail-intersimple-minobs.py +++ b/scratch/etienne/trpo/experiments/wgail-intersimple-minobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Minobs +from util.wrappers import CollisionPenaltyWrapper, Minobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/wgail-intersimple-minobs2.py b/scratch/etienne/trpo/experiments/wgail-intersimple-minobs2.py similarity index 97% rename from scratch/etienne/trpo/wgail-intersimple-minobs2.py rename to scratch/etienne/trpo/experiments/wgail-intersimple-minobs2.py index 733a95c..1b97041 100644 --- a/scratch/etienne/trpo/wgail-intersimple-minobs2.py +++ b/scratch/etienne/trpo/experiments/wgail-intersimple-minobs2.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Minobs +from util.wrappers import CollisionPenaltyWrapper, Minobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/wgail-intersimple-setobs2.py b/scratch/etienne/trpo/experiments/wgail-intersimple-setobs2.py similarity index 97% rename from scratch/etienne/trpo/wgail-intersimple-setobs2.py rename to scratch/etienne/trpo/experiments/wgail-intersimple-setobs2.py index 95b8ebc..08e653c 100644 --- a/scratch/etienne/trpo/wgail-intersimple-setobs2.py +++ b/scratch/etienne/trpo/experiments/wgail-intersimple-setobs2.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Setobs +from util.wrappers import CollisionPenaltyWrapper, Setobs import numpy as np from gym.wrappers import TransformObservation from torch.utils.tensorboard import SummaryWriter diff --git a/scratch/etienne/trpo/wgail-intersimple.py b/scratch/etienne/trpo/experiments/wgail-intersimple.py similarity index 96% rename from scratch/etienne/trpo/wgail-intersimple.py rename to scratch/etienne/trpo/experiments/wgail-intersimple.py index 793cb60..57d1f40 100644 --- a/scratch/etienne/trpo/wgail-intersimple.py +++ b/scratch/etienne/trpo/experiments/wgail-intersimple.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat( n_rays=5, diff --git a/scratch/etienne/trpo/wgail-options-setobs.py b/scratch/etienne/trpo/experiments/wgail-options-setobs.py similarity index 97% rename from scratch/etienne/trpo/wgail-options-setobs.py rename to scratch/etienne/trpo/experiments/wgail-options-setobs.py index 89944ed..eaf2d5b 100644 --- a/scratch/etienne/trpo/wgail-options-setobs.py +++ b/scratch/etienne/trpo/experiments/wgail-options-setobs.py @@ -9,7 +9,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter diff --git a/scratch/etienne/trpo/wgail-options-setobs2.py b/scratch/etienne/trpo/experiments/wgail-options-setobs2.py similarity index 96% rename from scratch/etienne/trpo/wgail-options-setobs2.py rename to scratch/etienne/trpo/experiments/wgail-options-setobs2.py index 6f7b2af..a851a46 100644 --- a/scratch/etienne/trpo/wgail-options-setobs2.py +++ b/scratch/etienne/trpo/experiments/wgail-options-setobs2.py @@ -9,7 +9,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter @@ -50,7 +50,7 @@ value = SetValue() v_opt = torch.optim.Adam(value.parameters(), lr=1e-4) discriminator = DeepsetDiscriminator() -disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-3) +disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-4) expert_data = torch.load('intersimple-expert-data-setobs2.pt') expert_data = Buffer(*expert_data) diff --git a/scratch/etienne/trpo/wgail-pendulum.py b/scratch/etienne/trpo/experiments/wgail-pendulum.py similarity index 100% rename from scratch/etienne/trpo/wgail-pendulum.py rename to scratch/etienne/trpo/experiments/wgail-pendulum.py diff --git a/scratch/etienne/trpo/wgail-ppo-intersimple-minobs.py b/scratch/etienne/trpo/experiments/wgail-ppo-intersimple-minobs.py similarity index 97% rename from scratch/etienne/trpo/wgail-ppo-intersimple-minobs.py rename to scratch/etienne/trpo/experiments/wgail-ppo-intersimple-minobs.py index 45d0143..ee9757d 100644 --- a/scratch/etienne/trpo/wgail-ppo-intersimple-minobs.py +++ b/scratch/etienne/trpo/experiments/wgail-ppo-intersimple-minobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Minobs +from util.wrappers import CollisionPenaltyWrapper, Minobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/wgail-ppo-intersimple-setobs2.py b/scratch/etienne/trpo/experiments/wgail-ppo-intersimple-setobs2.py similarity index 97% rename from scratch/etienne/trpo/wgail-ppo-intersimple-setobs2.py rename to scratch/etienne/trpo/experiments/wgail-ppo-intersimple-setobs2.py index 30b69a0..fe7472e 100644 --- a/scratch/etienne/trpo/wgail-ppo-intersimple-setobs2.py +++ b/scratch/etienne/trpo/experiments/wgail-ppo-intersimple-setobs2.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Setobs +from util.wrappers import CollisionPenaltyWrapper, Setobs import numpy as np from gym.wrappers import TransformObservation from torch.utils.tensorboard import SummaryWriter diff --git a/scratch/etienne/trpo/wgail-ppo-intersimple.py b/scratch/etienne/trpo/experiments/wgail-ppo-intersimple.py similarity index 100% rename from scratch/etienne/trpo/wgail-ppo-intersimple.py rename to scratch/etienne/trpo/experiments/wgail-ppo-intersimple.py diff --git a/scratch/etienne/trpo/wgail-ppo-options-setobs.py b/scratch/etienne/trpo/experiments/wgail-ppo-options-setobs.py similarity index 97% rename from scratch/etienne/trpo/wgail-ppo-options-setobs.py rename to scratch/etienne/trpo/experiments/wgail-ppo-options-setobs.py index fa89bf0..3f7ed6d 100644 --- a/scratch/etienne/trpo/wgail-ppo-options-setobs.py +++ b/scratch/etienne/trpo/experiments/wgail-ppo-options-setobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter diff --git a/scratch/etienne/trpo/wgail-ppo-options-setobs2.py b/scratch/etienne/trpo/experiments/wgail-ppo-options-setobs2.py similarity index 96% rename from scratch/etienne/trpo/wgail-ppo-options-setobs2.py rename to scratch/etienne/trpo/experiments/wgail-ppo-options-setobs2.py index ea6fbfa..d4394a2 100644 --- a/scratch/etienne/trpo/wgail-ppo-options-setobs2.py +++ b/scratch/etienne/trpo/experiments/wgail-ppo-options-setobs2.py @@ -1,4 +1,3 @@ -# %% import gym from options.options import gail_ppo, Buffer from core.value import SetValue @@ -8,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter @@ -51,12 +50,11 @@ value = SetValue() v_opt = torch.optim.Adam(value.parameters(), lr=1e-3) discriminator = DeepsetDiscriminator() -disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-3) +disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-4) expert_data = torch.load('intersimple-expert-data-setobs2.pt') expert_data = Buffer(*expert_data) -# %% value, policy = gail_ppo( env_fn=env_fn, expert_data=expert_data, diff --git a/scratch/etienne/trpo/wgail-ppo-pendulum.py b/scratch/etienne/trpo/experiments/wgail-ppo-pendulum.py similarity index 100% rename from scratch/etienne/trpo/wgail-ppo-pendulum.py rename to scratch/etienne/trpo/experiments/wgail-ppo-pendulum.py diff --git a/scratch/etienne/trpo/sb3/sb3-ppo-intersimple-rollout.py b/scratch/etienne/trpo/sb3/sb3-ppo-intersimple-rollout.py index eee73fe..2f9b6fd 100644 --- a/scratch/etienne/trpo/sb3/sb3-ppo-intersimple-rollout.py +++ b/scratch/etienne/trpo/sb3/sb3-ppo-intersimple-rollout.py @@ -7,7 +7,7 @@ from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools import torch -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper model = PPO.load('sb3-ppo-intersimple') env = CollisionPenaltyWrapper(IntersimpleLidarFlat( diff --git a/src/core/reparam_module.py b/src/core/reparam_module.py index 5bcd613..1b24986 100644 --- a/src/core/reparam_module.py +++ b/src/core/reparam_module.py @@ -160,3 +160,6 @@ class ReparamPolicy(ReparamModule): def predict(self, *args, **kwargs): return self.module.predict(*args, **kwargs) + + def unsafe_probability_mass(self, *args, **kwargs): + return self.module.unsafe_probability_mass(*args, **kwargs) diff --git a/src/options/envs.py b/src/options/envs.py new file mode 100644 index 0000000..6f6bdd4 --- /dev/null +++ b/src/options/envs.py @@ -0,0 +1,111 @@ +import gym +import numpy as np +from src.gail2.wrappers import Wrapper, Setobs, TransformObservation +from intersim.envs import IntersimpleLidarFlatIncrementingAgent + +obs_min = np.array([ + [-1000, -1000, 0, -np.pi, -1e-1, 0.], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], +]).reshape(-1) + +obs_max = np.array([ + [1000, 1000, 20, np.pi, 1e-1, 0.], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], +]).reshape(-1) + +def NormalizedOptionsEvalEnv(**kwargs): + return OptionsEnv(Setobs( + TransformObservation(IntersimpleLidarFlatIncrementingAgent( + n_rays=5, + **kwargs, + ), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) + ), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5)]) + +def NormalizedContinuousEvalEnv(**kwargs): + return Setobs( + TransformObservation(IntersimpleLidarFlatIncrementingAgent( + n_rays=5, + **kwargs, + ), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) + ) + +class OptionsEnv(Wrapper): + + def __init__(self, env, options): + super().__init__(env) + self.ll_action_space = env.action_space + self.options = options + self.action_space = gym.spaces.Discrete(len(options)) + self.max_plan_length = max(t for _, t in options) + + def plan(self, option): + target_v, t = option + current_v = self.env._env.state[self.env._agent, 1].item() + dt = self.env._env._dt + a = (target_v - current_v) / (t * dt) + a = self.env._normalize(a) + a = a * np.ones((t,)) + a += 0.01 * np.random.randn(*a.shape) + a = np.clip(a, self.ll_action_space.low, self.ll_action_space.high) + return a + + def execute_plan(self, obs, option, render_mode=None): + observations = np.zeros((self.max_plan_length + 1, *self.env.observation_space.shape)) + actions = np.zeros((self.max_plan_length + 1, *self.ll_action_space.shape)) + rewards = np.zeros((self.max_plan_length + 1,)) + env_done = np.ones((self.max_plan_length + 1,), dtype=bool) + plan_done = np.ones((self.max_plan_length + 1,), dtype=bool) + infos = [] + + observations[0] = obs + env_done[0] = False + for k, u in enumerate(self.plan(option)): + plan_done[k] = False + o, r, d, i = super().step(u) + actions[k] = u + rewards[k] = r + env_done[k] = d + infos.append(i) + observations[k+1] = o + + if render_mode is not None: + self.env.render(render_mode) + + if d: + break + + n_steps = k + 1 + return observations, actions, rewards, env_done, plan_done, infos, n_steps + + def step(self, action, render_mode=None): + a = int(action) + assert a == action + ll_obs, ll_actions, ll_rewards, ll_env_done, ll_plan_done, ll_infos, ll_steps = self.execute_plan(self.last_obs, self.options[a], render_mode) + hl_obs = ll_obs[ll_steps] + hl_reward = (ll_rewards * ~ll_plan_done).sum().item() + hl_done = ll_env_done[ll_steps-1].item() + hl_infos = { + 'll': { + 'observations': ll_obs, + 'actions': ll_actions, + 'rewards': ll_rewards, + 'env_done': ll_env_done, + 'plan_done': ll_plan_done, + 'infos': ll_infos, + 'steps': ll_steps, + } + } + self.last_obs = hl_obs + return hl_obs, hl_reward, hl_done, hl_infos + + def reset(self, *args, **kwargs): + self.last_obs = super().reset(*args, **kwargs) + return self.last_obs diff --git a/scratch/etienne/trpo/options/options.py b/src/options/options.py similarity index 96% rename from scratch/etienne/trpo/options/options.py rename to src/options/options.py index af0414a..0453fd5 100644 --- a/scratch/etienne/trpo/options/options.py +++ b/src/options/options.py @@ -18,7 +18,7 @@ class OptionsRollout: def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, - gae_lambda, delta, backtrack_coeff, backtrack_iters, cg_iters=10, cg_damping=0.1, wasserstein=False, wasserstein_c=None, logger=TerminalLogger()): + gae_lambda, delta, backtrack_coeff, backtrack_iters, cg_iters=10, cg_damping=0.1, wasserstein=False, wasserstein_c=None, logger=TerminalLogger(), callback=None): policy(torch.zeros(env_fn(0).observation_space.shape)) policy = ReparamPolicy(policy) @@ -49,11 +49,14 @@ def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value value, policy = trpo_step(value, policy, generator_data.hl.states, generator_data.hl.actions, generator_data.hl.rewards, generator_data.hl.dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters, cg_damping) expert_data = roll_buffer(expert_data, shifts=-3, dims=0) + if callback is not None: + callback(epoch, value, policy) + return value, policy def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, - gae_lambda, clip_ratio, pi_opt, pi_iters, target_kl=None, max_grad_norm=None, wasserstein=False, wasserstein_c=None, logger=TerminalLogger()): + gae_lambda, clip_ratio, pi_opt, pi_iters, target_kl=None, max_grad_norm=None, wasserstein=False, wasserstein_c=None, logger=TerminalLogger(), callback=None): logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0]) logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0]) @@ -81,6 +84,9 @@ def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, v value, policy = ppo_step(value, policy, generator_data.hl.states, generator_data.hl.actions, generator_data.hl.rewards, generator_data.hl.dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm) expert_data = roll_buffer(expert_data, shifts=-3, dims=0) + if callback is not None: + callback(epoch, value, policy) + return value, policy def rollout(env_fn, policy, n_episodes, max_steps_per_episode): @@ -99,7 +105,6 @@ def rollout(env_fn, policy, n_episodes, max_steps_per_episode): env = VecEnv(list(map(lambda i: (lambda: env_fn(i)), range(n_episodes)))) states[:, 0] = torch.tensor(env.reset()).clone().detach() - dones[:, 0] = False for s in tqdm(range(max_steps_per_episode), 'Rollout'): actions[:, s] = policy.sample(policy(states[:, s])).clone().detach() @@ -111,7 +116,7 @@ def rollout(env_fn, policy, n_episodes, max_steps_per_episode): o, r, d, info = env.step(clipped_actions) states[:, s + 1] = torch.tensor(o).clone().detach() rewards[:, s] = torch.tensor(r).clone().detach() - dones[:, s + 1] = torch.tensor(d).clone().detach() + dones[:, s] = torch.tensor(d).clone().detach() ll_states[:, s] = torch.from_numpy(np.stack([i['ll']['observations'] for i in info])).clone().detach() ll_actions[:, s] = torch.from_numpy(np.stack([i['ll']['actions'] for i in info])).clone().detach() @@ -162,7 +167,7 @@ class OptionsEnv(gym.Wrapper): o, r, d, i = super().step(u) actions[k] = u rewards[k] = r - env_done[k+1] = d + env_done[k] = d infos.append(i) observations[k+1] = o @@ -181,7 +186,7 @@ class OptionsEnv(gym.Wrapper): ll_obs, ll_actions, ll_rewards, ll_env_done, ll_plan_done, ll_infos, ll_steps = self.execute_plan(self.last_obs, self.options[a], render_mode) hl_obs = ll_obs[ll_steps] hl_reward = (ll_rewards * ~ll_plan_done).sum().item() - hl_done = ll_env_done[ll_steps].item() + hl_done = ll_env_done[ll_steps-1].item() hl_infos = { 'll': { 'observations': ll_obs, diff --git a/scratch/etienne/trpo/options/test_options.py b/src/options/test_options.py similarity index 100% rename from scratch/etienne/trpo/options/test_options.py rename to src/options/test_options.py diff --git a/src/safe_options/collisions.py b/src/safe_options/collisions.py new file mode 100644 index 0000000..36a6448 --- /dev/null +++ b/src/safe_options/collisions.py @@ -0,0 +1,185 @@ +import torch +import numpy as np +from intersim.collisions import state_to_polygon + +def safety_plan(env, plan): + return np.concatenate((plan, np.array(5 * [env._env._min_acc])), axis=0) + +def available_actions(env, options): + """Return mask of available actions given current `env` state.""" + plans = [generate_plan(env, i, options) for i, _ in enumerate(options)] + # is emergency braking still possible? + plans = list(map(lambda p: safety_plan(env, p), plans)) + + T = max(len(p) for p in plans) + plans = [np.pad(p, ((0, T-len(p)),), constant_values=np.nan) for p in plans] + plans = np.stack(plans, axis=0) + + valid = feasible(env, plans) + return valid + +def target_velocity_plan(current_v: float, target_v: float, t: int, dt: float): + """Smoothly target a velocity in a given number of steps""" + # for now, constant acceleration + a = (target_v - current_v) / (t * dt) + return a*np.ones((t,)) + +def generate_plan(env, i, options): + """Generate input profile for high-level action `i`.""" + assert i < len(options), "Invalid option index {i}" + target_v, t = options[i] + current_v = env._env.state[env._agent, 1].item() # extract from env + plan = target_velocity_plan(current_v, target_v, t, env._env._dt) + assert len(plan) == t, "incorrect plan length" + return plan + +def feasible(env, plan, method='exact'): + """Check if input profile is feasible given current `env` state.""" + # zero pad plan - Take (B, T) or (T,) np plan and convert it to (B, T, nv, 1) torch.Tensor + plan = torch.tensor(plan) + plan = plan.reshape(-1, plan.shape[-1]) + full_plan = torch.zeros(*plan.shape, env._env._nv, 1) + full_plan[:, :, env._agent, 0] = plan + + # check_future_collisions_fast takes in B-list and outputs (B,) bool tensor + if method=='circle': + valid = check_future_collisions_fast(env, full_plan) + elif method=='ncircles': + valid = check_future_collisions_ncircles(env, full_plan) + elif method=='exact': + valid = check_future_collisions_exact(env, full_plan) + else: + raise NotImplementedError('Invalid collision-checking method') + + return valid + +def check_future_collisions_ncircles(env, actions, n_circles:int=2): + """Checks whether `env._agent` would collide with other agents assuming `actions` as input. + + Vehicles are (over-)approximated by multiple circles. + + Args: + env (gym.Env): current environment state + actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles + Returns: + feasible (torch.Tensor): tensor of shape (B,) indicating whether the respective action profiles are collision-free + """ + assert n_circles >= 2 + B, (T, nv, _) = len(actions), actions[0].shape + + states = env._env.propagate_action_profile_vectorized(actions) + assert states.shape == (B, T, nv, 5) + centers = states[:, :, :, :2] + psi = states[:, :, :, 3] + lon = torch.stack([psi.cos(), psi.sin()],dim=-1) # (B, T, nv, 2) + + # offset between [-env._env.lengths+env._env.widths/2, env._env.lengths/2-env._env.widths/2] + back = (-env._env._lengths/2+env._env._widths/2).unsqueeze(-1) # (nv, 1) + length = (env._env._lengths-env._env._widths).unsqueeze(-1) # (nv, 1) + diff_d = back + length*(torch.arange(n_circles)/(n_circles-1)).unsqueeze(0) # (nv, n_circles) + assert diff_d.shape == (nv, n_circles) + + offsets = diff_d[None, None, :, :, None] * lon[:, :, :, None, :] + assert offsets.shape == (B, T, nv, n_circles, 2) + + expanded_centers=centers.unsqueeze(-2) + offsets #(B, T, nv, n_circles, 2) + assert expanded_centers.shape == (B, T, nv, n_circles, 2) + agent_centers = expanded_centers[:,:,env._agent:env._agent+1,:,:] #(B, T, 1, n_circles, 2) + ds = expanded_centers.reshape((B, T, nv*n_circles, 1, 2)) - agent_centers #(B, T, nv*nc,1, 2) - (B, T, 1, nc, 2) = (B, T, nv*nc, nc, 2) + + distance = (ds**2).sum(-1).sqrt().reshape((B, T, nv, n_circles, n_circles)) # (B, T, nv, nc, nc) + distance = torch.where(distance.isnan(), np.inf*torch.ones_like(distance), distance) # only collide with spawned agents + distance[:, :, env._agent] = np.inf # cannot collide with itself + assert distance.shape == (B, T, nv, n_circles, n_circles) + + radius = env._env._widths*np.sqrt(2) / 2 + min_distance = radius[env._agent] + radius + min_distance = min_distance[None, None, :, None, None] + assert min_distance.shape == (1, 1, nv, 1, 1) + + return (distance > min_distance).all(-1).all(-1).all(-1).all(-1) + +def check_future_collisions_circle(env, actions): + """Compute collision information for circular vehicle approximations + + Args: + env (gym.Env): current environment state + actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles + Returns: + states (torch.Tensor): tensor of shape (B, T, nv, 5) of future states based on the action profiles + collision_tensor (torch.Tensor): tensor of shape (B, T, nv) of bools indicating which plan collides with which vehicles in which time frame + false: colliding, true: not colliding + """ + B, (T, nv, _) = len(actions), actions[0].shape + + states = env._env.propagate_action_profile_vectorized(actions) + assert states.shape == (B, T, nv, 5) + + distance = ((states[:, :, :, :2] - states[:, :, env._agent:env._agent+1, :2])**2).sum(-1).sqrt() + distance = torch.where(distance.isnan(), np.inf*torch.ones_like(distance), distance) # only collide with spawned agents + distance[:, :, env._agent] = np.inf # cannot collide with itself + assert distance.shape == (B, T, nv) + + radius = (env._env._lengths**2 + env._env._widths**2).sqrt() / 2 + min_distance = radius[env._agent] + radius + min_distance = min_distance.unsqueeze(0).unsqueeze(0) + assert min_distance.shape == (1, 1, nv) + + collision_tensor = distance > min_distance + assert collision_tensor.shape == (B, T, nv) + return states, collision_tensor + +def check_future_collisions_fast(env, actions): + """Checks whether `env._agent` would collide with other agents assuming `actions` as input. + + Vehicles are (over-)approximated by single circles. + + Args: + env (gym.Env): current environment state + actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles + Returns: + feasible (torch.Tensor): tensor of shape (B,) indicating whether the respective action profiles are collision-free + """ + _, collision_tensor = check_future_collisions_circle(env, actions) + return collision_tensor.all(-1).all(-1) + +def check_future_collisions_exact(env, actions): + """ + Checks whether `env._agent` would collide with other agents assuming `actions` as input. + + Args: + env (gym.Env): current environment state + actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles + Returns: + feasible (torch.Tensor): tensor of shape (B,) indicating whether the respective action profiles are collision-free + """ + # First check with simple circle collision check + states, collision_tensor = check_future_collisions_circle(env, actions) + (B, T, nv, _) = states.shape + # For those that have colliding circles, check exactly + colliding_mask = ~collision_tensor + + ego_states = states[:, :, env._agent:env._agent+1, :].expand(states.shape) + assert ego_states.shape == states.shape + + # get dimensions + lengths = env._env._lengths.expand(states.shape[:3]) + widths = env._env._widths.expand(states.shape[:3]) + ego_lengths = lengths[:, :, env._agent:env._agent+1].expand(lengths.shape) + ego_widths = widths[:, :, env._agent:env._agent+1].expand(widths.shape) + assert lengths.shape == widths.shape == ego_lengths.shape == ego_widths.shape == (B, T, nv) + + # For every collision instance between ego and other vehicle, check whether rectangles intersect + exact_collisions = torch.zeros_like(collision_tensor[colliding_mask]) + for i, (ego_state, ego_length, ego_width, other_state, other_length, other_width) in enumerate(zip( + ego_states[colliding_mask], ego_lengths[colliding_mask], ego_widths[colliding_mask], + states[colliding_mask], lengths[colliding_mask], widths[colliding_mask] + )): + assert ego_state.shape == other_state.shape == (5,) + assert ego_length.shape == ego_width.shape == other_length.shape == other_width.shape == () + p_ego = state_to_polygon(ego_state, ego_length, ego_width) + p_other = state_to_polygon(other_state, other_length, other_width) + exact_collisions[i] = p_ego.intersects(p_other) + + collision_tensor[colliding_mask] = ~exact_collisions + return collision_tensor.all(-1).all(-1) diff --git a/src/safe_options/options.py b/src/safe_options/options.py new file mode 100644 index 0000000..a8d0e50 --- /dev/null +++ b/src/safe_options/options.py @@ -0,0 +1,305 @@ +import gym +import numpy as np +import torch +from stable_baselines3.common.vec_env import DummyVecEnv as VecEnv + +from core.reparam_module import ReparamPolicy +from tqdm import tqdm +from core.gail import train_discriminator, roll_buffer, TerminalLogger +from dataclasses import dataclass +from safe_options.policy_gradient import trpo_step, ppo_step +import torch.nn.functional as F + +from safe_options.collisions import feasible + +@dataclass +class Buffer: + states: torch.Tensor + actions: torch.Tensor + rewards: torch.Tensor + dones: torch.Tensor + +@dataclass +class HLBuffer: + states: torch.Tensor + safe_actions: torch.Tensor + actions: torch.Tensor + rewards: torch.Tensor + dones: torch.Tensor + +@dataclass +class OptionsRollout: + hl: HLBuffer + ll: Buffer + +def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, + v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, + gae_lambda, delta, backtrack_coeff, backtrack_iters, cg_iters=10, cg_damping=0.1, wasserstein=False, wasserstein_c=None, logger=TerminalLogger(), callback=None): + + policy(torch.zeros(env_fn(0).observation_space['observation'].shape), torch.zeros(env_fn(0).observation_space['safe_actions'].shape)) + policy = ReparamPolicy(policy) + + logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0]) + logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0]) + + for epoch in tqdm(range(epochs)): + hl_data, ll_data = rollout(env_fn, policy, rollout_episodes, rollout_steps) + generator_data = OptionsRollout(HLBuffer(*hl_data), Buffer(*ll_data)) + + generator_data.ll.actions += 0.1 * torch.randn_like(generator_data.ll.actions) + + logger.add_scalar('gen/mean_episode_length', (~generator_data.ll.dones).sum() / generator_data.ll.states.shape[0], epoch) + logger.add_scalar('gen/mean_reward_per_episode', generator_data.hl.rewards[~generator_data.hl.dones].sum() / generator_data.hl.states.shape[0], epoch) + logger.add_scalar('gen/unsafe_probability_mass', policy.unsafe_probability_mass(policy(generator_data.hl.states[~generator_data.hl.dones], generator_data.hl.safe_actions[~generator_data.hl.dones])).mean(), epoch) + + discriminator, loss = train_discriminator(expert_data, generator_data.ll, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c) + if wasserstein: + generator_data.ll.rewards = discriminator(generator_data.ll.states, generator_data.ll.actions) + else: + generator_data.ll.rewards = -F.logsigmoid(discriminator(generator_data.ll.states, generator_data.ll.actions)) + logger.add_scalar('disc/final_loss', loss, epoch) + logger.add_scalar('disc/mean_reward_per_episode', generator_data.ll.rewards[~generator_data.ll.dones].sum() / generator_data.ll.states.shape[0], epoch) + + #assert generator_data.ll.rewards.shape == generator_data.ll.dones.shape + generator_data.hl.rewards = torch.where(~generator_data.ll.dones, generator_data.ll.rewards, torch.tensor(0.)).sum(-1) + + value, policy = trpo_step(value, policy, generator_data.hl.states, generator_data.hl.safe_actions, generator_data.hl.actions, generator_data.hl.rewards, generator_data.hl.dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters, cg_damping) + expert_data = roll_buffer(expert_data, shifts=-3, dims=0) + + if callback is not None: + callback(epoch, value, policy) + + return value, policy + +def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, + v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, + gae_lambda, clip_ratio, pi_opt, pi_iters, target_kl=None, max_grad_norm=None, wasserstein=False, wasserstein_c=None, logger=TerminalLogger(), callback=None): + + logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0]) + logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0]) + + for epoch in range(epochs): + hl_data, ll_data = rollout(env_fn, policy, rollout_episodes, rollout_steps) + generator_data = OptionsRollout(HLBuffer(*hl_data), Buffer(*ll_data)) + + generator_data.ll.actions += 0.1 * torch.randn_like(generator_data.ll.actions) + + logger.add_scalar('gen/mean_episode_length', (~generator_data.ll.dones).sum() / generator_data.ll.states.shape[0], epoch) + logger.add_scalar('gen/mean_reward_per_episode', generator_data.hl.rewards[~generator_data.hl.dones].sum() / generator_data.hl.states.shape[0], epoch) + logger.add_scalar('gen/unsafe_probability_mass', policy.unsafe_probability_mass(policy(generator_data.hl.states[~generator_data.hl.dones], generator_data.hl.safe_actions[~generator_data.hl.dones])).mean(), epoch) + + discriminator, loss = train_discriminator(expert_data, generator_data.ll, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c) + if wasserstein: + generator_data.ll.rewards = discriminator(generator_data.ll.states, generator_data.ll.actions) + else: + generator_data.ll.rewards = -F.logsigmoid(discriminator(generator_data.ll.states, generator_data.ll.actions)) + logger.add_scalar('disc/final_loss', loss, epoch) + logger.add_scalar('disc/mean_reward_per_episode', generator_data.ll.rewards[~generator_data.ll.dones].sum() / generator_data.ll.states.shape[0], epoch) + + #assert generator_data.ll.rewards.shape == generator_data.ll.dones.shape + generator_data.hl.rewards = torch.where(~generator_data.ll.dones, generator_data.ll.rewards, torch.tensor(0.)).sum(-1) + + value, policy = ppo_step(value, policy, generator_data.hl.states, generator_data.hl.safe_actions, generator_data.hl.actions, generator_data.hl.rewards, generator_data.hl.dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm) + expert_data = roll_buffer(expert_data, shifts=-3, dims=0) + + if callback is not None: + callback(epoch, value, policy) + + return value, policy + +def rollout(env_fn, policy, n_episodes, max_steps_per_episode): + env = env_fn(0) + + states = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.observation_space['observation'].shape) + safe_actions = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.observation_space['safe_actions'].shape) + actions = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.action_space.shape) + rewards = torch.zeros(n_episodes, max_steps_per_episode + 1) + dones = torch.ones(n_episodes, max_steps_per_episode + 1, dtype=bool) + + ll_states = torch.zeros(n_episodes, max_steps_per_episode, env.max_plan_length + 1, *env.observation_space['observation'].shape) + ll_actions = torch.zeros(n_episodes, max_steps_per_episode, env.max_plan_length + 1, *env.ll_action_space.shape) + ll_rewards = torch.zeros(n_episodes, max_steps_per_episode, env.max_plan_length + 1) + ll_dones = torch.ones(n_episodes, max_steps_per_episode, env.max_plan_length + 1, dtype=bool) + + env = VecEnv(list(map(lambda i: (lambda: env_fn(i)), range(n_episodes)))) + + obs = env.reset() + states[:, 0] = torch.tensor(obs['observation']).clone().detach() + safe_actions[:, 0] = torch.tensor(obs['safe_actions']).clone().detach() + dones[:, 0] = False + + for s in tqdm(range(max_steps_per_episode), 'Rollout'): + actions[:, s] = policy.sample(policy(states[:, s], safe_actions[:, s])).clone().detach() + + clipped_actions = actions[:, s] + if isinstance(env.action_space, gym.spaces.Box): + clipped_actions = torch.clamp(clipped_actions, torch.from_numpy(env.action_space.low), torch.from_numpy(env.action_space.high)) + + o, r, d, info = env.step(clipped_actions) + states[:, s + 1] = torch.tensor(o['observation']).clone().detach() + safe_actions[:, s + 1] = torch.tensor(o['safe_actions']).clone().detach() + rewards[:, s] = torch.tensor(r).clone().detach() + dones[:, s + 1] = torch.tensor(d).clone().detach() + + ll_states[:, s] = torch.from_numpy(np.stack([i['ll']['observations'] for i in info])).clone().detach() + ll_actions[:, s] = torch.from_numpy(np.stack([i['ll']['actions'] for i in info])).clone().detach() + ll_rewards[:, s] = torch.from_numpy(np.stack([i['ll']['rewards'] for i in info])).clone().detach() + ll_dones[:, s] = torch.from_numpy(np.stack([i['ll']['plan_done'] for i in info])).clone().detach() + + dones = dones.cumsum(1) > 0 + + states = states[:, :max_steps_per_episode] + safe_actions = safe_actions[:, :max_steps_per_episode] + actions = actions[:, :max_steps_per_episode] + rewards = rewards[:, :max_steps_per_episode] + dones = dones[:, :max_steps_per_episode] + + return (states, safe_actions, actions, rewards, dones), (ll_states, ll_actions, ll_rewards, ll_dones) + +class OptionsEnv(gym.Wrapper): + + def __init__(self, env, options): + super().__init__(env) + self.ll_action_space = env.action_space + self.options = options + self.action_space = gym.spaces.Discrete(len(options)) + self.max_plan_length = max(t for _, t in options) + + def plan(self, option): + target_v, t = option + current_v = self.env._env.state[self.env._agent, 1].item() + dt = self.env._env._dt + a = (target_v - current_v) / (t * dt) + a = self.env._normalize(a) + a = a * np.ones((t,)) + a += 0.01 * np.random.randn(*a.shape) + a = np.clip(a, self.ll_action_space.low, self.ll_action_space.high) + return a + + def execute_plan(self, obs, option, render_mode=None): + observations = np.zeros((self.max_plan_length + 1, *self.env.observation_space.shape)) + actions = np.zeros((self.max_plan_length + 1, *self.ll_action_space.shape)) + rewards = np.zeros((self.max_plan_length + 1,)) + env_done = np.ones((self.max_plan_length + 1,), dtype=bool) + plan_done = np.ones((self.max_plan_length + 1,), dtype=bool) + infos = [] + + plan = self.plan(option) + observations[0] = obs + env_done[0] = False + for k, u in enumerate(plan): + plan_done[k] = False + o, r, d, i = self.env.step(u) + actions[k] = u + rewards[k] = r + env_done[k+1] = d + infos.append(i) + observations[k+1] = o + + if render_mode is not None: + self.env.render(render_mode) + + if d: + break + + n_steps = k + 1 + return observations, actions, rewards, env_done, plan_done, infos, n_steps + + def step(self, action, render_mode=None): + a = int(action) + assert a == action + ll_obs, ll_actions, ll_rewards, ll_env_done, ll_plan_done, ll_infos, ll_steps = self.execute_plan(self.last_obs, self.options[a], render_mode) + hl_obs = ll_obs[ll_steps] + hl_reward = (ll_rewards * ~ll_plan_done).sum().item() + hl_done = ll_env_done[ll_steps].item() + hl_infos = { + 'll': { + 'observations': ll_obs, + 'actions': ll_actions, + 'rewards': ll_rewards, + 'env_done': ll_env_done, + 'plan_done': ll_plan_done, + 'infos': ll_infos, + 'steps': ll_steps, + } + } + self.last_obs = hl_obs + return hl_obs, hl_reward, hl_done, hl_infos + + def reset(self, *args, **kwargs): + self.last_obs = super().reset(*args, **kwargs) + return self.last_obs + + +class SafeOptionsEnv(OptionsEnv): + + def __init__(self, env, options, safe_actions_collision_method=None, abort_unsafe_collision_method=None): + super().__init__(env, options) + self.safe_actions_collision_method = safe_actions_collision_method + self.abort_unsafe_collision_method = abort_unsafe_collision_method + self.observation_space = gym.spaces.Dict({ + 'observation': self.observation_space, + 'safe_actions': gym.spaces.Box(low=0., high=1., shape=(self.action_space.n,)), + }) + + def safe_actions(self): + if self.safe_actions_collision_method is None: + return np.ones(len(self.options), dtype=bool) + + plans = [self.plan(o) for o in self.options] + plans = np.stack(plans) + safe = feasible(self.env, plans, method=self.safe_actions_collision_method) + if not safe.any(): + # action 0 is considered safe fallback + safe[0] = True + + return safe + + def reset(self, *args, **kwargs): + obs = super().reset(*args, **kwargs) + obs = { + 'observation': obs, + 'safe_actions': self.safe_actions(), + } + return obs + + def step(self, action, render_mode=None): + obs, reward, done, info = super().step(action, render_mode) + obs = { + 'observation': obs, + 'safe_actions': self.safe_actions(), + } + return obs, reward, done, info + + def execute_plan(self, obs, option, render_mode=None): + observations = np.zeros((self.max_plan_length + 1, *self.env.observation_space.shape)) + actions = np.zeros((self.max_plan_length + 1, *self.ll_action_space.shape)) + rewards = np.zeros((self.max_plan_length + 1,)) + env_done = np.ones((self.max_plan_length + 1,), dtype=bool) + plan_done = np.ones((self.max_plan_length + 1,), dtype=bool) + infos = [] + + plan = self.plan(option) + observations[0] = obs + env_done[0] = False + for k, u in enumerate(plan): + plan_done[k] = False + o, r, d, i = self.env.step(u) + actions[k] = u + rewards[k] = r + env_done[k+1] = d + infos.append(i) + observations[k+1] = o + + if render_mode is not None: + self.env.render(render_mode) + + if d: + break + + if self.abort_unsafe_collision_method is not None and \ + not feasible(self.env, plan[k:], method=self.abort_unsafe_collision_method): + break + + n_steps = k + 1 + return observations, actions, rewards, env_done, plan_done, infos, n_steps diff --git a/src/safe_options/policy.py b/src/safe_options/policy.py new file mode 100644 index 0000000..4e0e9d0 --- /dev/null +++ b/src/safe_options/policy.py @@ -0,0 +1,35 @@ +import torch +import torch.nn as nn +from torch.distributions import Categorical +from torch.distributions.kl import kl_divergence +from core.policy import SetDiscretePolicy + +class SetMaskedDiscretePolicy(SetDiscretePolicy): + + def forward(self, observation, safe_actions): + return torch.cat((super().forward(observation), safe_actions), -1) + + def torch_dist(self, dist): + logits = dist[..., :self.action_dim] + z = dist[..., self.action_dim:] + a = super().torch_dist(logits).probs + return Categorical(probs=a*z) + + def unsafe_probability_mass(self, dist): + logits = dist[..., :self.action_dim] + z = dist[..., self.action_dim:] + a = super().torch_dist(logits).probs + return (a * (1 - z)).sum(-1) + + # def torch_dist_nomask(self, dist): + # print('no mask logprob') + # logits = dist[..., :self.action_dim] + # return super().torch_dist(logits) + + # def log_prob(self, dist, actions): + # return self.torch_dist_nomask(dist).log_prob(actions) + + # def kl_divergence(self, dist1, dist2): + # d1 = self.torch_dist_nomask(dist1) + # d2 = self.torch_dist_nomask(dist2) + # return kl_divergence(d1, d2) diff --git a/src/safe_options/policy_gradient.py b/src/safe_options/policy_gradient.py new file mode 100644 index 0000000..6e968e4 --- /dev/null +++ b/src/safe_options/policy_gradient.py @@ -0,0 +1,107 @@ +import torch +from core.value_estimation import gae +from core.optimization import conjugate_gradient, line_search + +def trpo_step(value, policy, states, safe_actions, actions, rewards, dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1): + + states = states.detach() + actions = actions.detach() + rewards = rewards.detach() + dones = dones.detach() + + advantages, returns, valid = gae(states, rewards, value(states), dones, gamma, gae_lambda) + advantages = advantages.detach() + returns = returns.detach() + + # update value function + + for _ in range(v_iters): + v_opt.zero_grad() + value_loss = (value(states) - returns).pow(2)[valid].mean() + value_loss.backward() + v_opt.step() + + # compute policy gradient + + plogprob = policy.log_prob(policy(states, safe_actions), actions) + surrogate_advantage = (plogprob * advantages)[valid].sum() / states.shape[0] + g = torch.cat(torch.autograd.grad(surrogate_advantage, policy.flat_param)).detach() + + def Hx(x): + kl = policy.kl_divergence(policy(states, safe_actions), policy(states, safe_actions).detach())[valid].mean() + dKL = torch.cat(torch.autograd.grad(kl, policy.flat_param, create_graph=True)) + H_x = torch.cat(torch.autograd.grad(dKL.T @ x, policy.flat_param)).detach() + return H_x + cg_damping * x + + x = conjugate_gradient(Hx, g, cg_iters) + npg = torch.sqrt(2 * delta / (x.T @ Hx(x))) * x + + # perform line search + + def L(theta): + rplogprob = policy.log_prob(policy(states, safe_actions, flat_param=theta), actions) + return ((rplogprob - plogprob.detach()).exp() * advantages)[valid].sum() / advantages.shape[0] + + condition = lambda theta: policy.kl_divergence(policy(states, safe_actions, flat_param=theta), policy(states, safe_actions))[valid].mean() < delta + + x0 = policy.flat_param + g0 = torch.cat(torch.autograd.grad(L(x0), x0)) + theta = line_search(L, x0, npg, g0, backtrack_coeff, condition, max_steps=backtrack_iters) + + # update policy parameters + + with torch.no_grad(): + policy.flat_param.copy_(theta) + + return value, policy + +def ppo_step(value, policy, states, safe_actions, actions, rewards, dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm): + + states = states.detach() + actions = actions.detach() + rewards = rewards.detach() + dones = dones.detach() + + advantages, returns, valid = gae(states, rewards, value(states), dones, gamma, gae_lambda) + advantages = advantages.detach() + returns = returns.detach() + + # update value function + + for _ in range(v_iters): + v_opt.zero_grad() + value_loss = (value(states) - returns).pow(2)[valid].mean() + value_loss.backward() + v_opt.step() + + # update policy + + old_dist = policy(states, safe_actions).detach() + old_logprob = policy.log_prob(old_dist, actions).detach() + + def g(advantages, clip_ratio): + return torch.where(advantages >= 0, (1 + clip_ratio) * advantages, (1 - clip_ratio) * advantages) + + def L(states, actions, advantages, clip_ratio): + return torch.minimum( + (policy.log_prob(policy(states, safe_actions), actions) - old_logprob).exp() * advantages, + g(advantages, clip_ratio) + )[valid].mean() + + for _ in range(pi_iters): + pi_opt.zero_grad() + ppo_loss = -L(states, actions, advantages, clip_ratio) + ppo_loss.backward() + + if max_grad_norm: + torch.nn.utils.clip_grad_norm(policy.parameters(), max_grad_norm) + + pi_opt.step() + + kl = policy.kl_divergence(policy(states, safe_actions), old_dist)[valid].mean() + if target_kl and kl > target_kl: + break + + print('KL', kl.item()) + + return value, policy diff --git a/src/safe_options/test_options.py b/src/safe_options/test_options.py new file mode 100644 index 0000000..269eb43 --- /dev/null +++ b/src/safe_options/test_options.py @@ -0,0 +1,54 @@ +from intersim.envs import IntersimpleLidarFlat +from options import OptionsEnv +import gym +import numpy as np + +def test_obs_shape(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + assert env.reset().shape == (36,) + +def test_act_space(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + assert env.action_space == gym.spaces.Discrete(3) + +def test_plan(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + env.reset() + plan = env.plan(options[0]) + assert np.allclose(plan, -13.998268127441406 * np.ones((5,))) + +def test_plan2(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + obs = env.reset() + states, actions, rewards, dones, plan_done, infos, n_steps = env.execute_plan(obs, options[0]) + assert states.shape == (6, 36) + assert rewards.shape == (6,) + assert dones.shape == (6,) + assert len(infos) == 5 + +def test_step(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + env.reset() + obs, reward, done, _ = env.step(0) + assert obs.shape == (36,) + assert reward == 5.0 + assert done == False + +def test_ll_step(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + env.reset() + _, _, _, info = env.step(0) + assert info['ll']['observations'].shape == (6, 36) + assert info['ll']['actions'].shape == (6, 1) + assert info['ll']['rewards'].shape == (6,) + assert info['ll']['env_done'].shape == (6,) + assert info['ll']['plan_done'].shape == (6,) + assert info['ll']['plan_done'][5] == True + assert info['ll']['steps'] == 5 + assert len(info['ll']['infos']) == 5 diff --git a/scratch/etienne/trpo/wrappers.py b/src/util/wrappers.py similarity index 100% rename from scratch/etienne/trpo/wrappers.py rename to src/util/wrappers.py