Merge updated files

This commit is contained in:
ebuehrle
2022-02-17 22:41:55 +01:00
parent b78f95bab5
commit 5bd8b42d9f
83 changed files with 1439 additions and 1123 deletions

View File

@@ -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

View File

@@ -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)

View File

@@ -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

View File

@@ -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

View File

@@ -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))

View File

@@ -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

View File

@@ -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)

View File

@@ -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

View File

@@ -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)

View File

@@ -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

View File

@@ -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)

View File

@@ -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

View File

@@ -26,7 +26,7 @@ torch.save(policy.state_dict(), 'bc-intersimple-setobs2.pt')
# %% # %%
import numpy as np import numpy as np
from core.policy import SetPolicy 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 import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Minobs from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -8,7 +8,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Minobs from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Setobs from util.wrappers import CollisionPenaltyWrapper, Setobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -8,7 +8,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Setobs from util.wrappers import CollisionPenaltyWrapper, Setobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy

View File

@@ -8,7 +8,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Setobs from util.wrappers import CollisionPenaltyWrapper, Setobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat( envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat(
n_rays=5, n_rays=5,

View File

@@ -8,7 +8,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter

View File

@@ -8,7 +8,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter

View File

@@ -9,7 +9,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter
@@ -56,6 +56,11 @@ expert_data = torch.load('intersimple-expert-data-setobs2.pt')
expert_data = Buffer(*expert_data) 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( value, policy = gail(
env_fn=env_fn, env_fn=env_fn,
expert_data=expert_data, expert_data=expert_data,
@@ -66,7 +71,7 @@ value, policy = gail(
value=value, value=value,
v_opt=v_opt, v_opt=v_opt,
v_iters=1000, v_iters=1000,
epochs=200, epochs=300,
rollout_episodes=60, rollout_episodes=60,
rollout_steps=60, rollout_steps=60,
gamma=0.99, gamma=0.99,
@@ -75,6 +80,7 @@ value, policy = gail(
backtrack_coeff=0.8, backtrack_coeff=0.8,
backtrack_iters=10, backtrack_iters=10,
logger=SummaryWriter(comment='gail-options-setobs2'), logger=SummaryWriter(comment='gail-options-setobs2'),
callback=callback,
) )
torch.save(policy.state_dict(), 'gail-options-setobs2.pt') torch.save(policy.state_dict(), 'gail-options-setobs2.pt')

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Minobs from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -8,7 +8,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Setobs from util.wrappers import CollisionPenaltyWrapper, Setobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat( envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat(
n_rays=5, n_rays=5,

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter

View File

@@ -8,7 +8,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter
@@ -57,6 +57,11 @@ expert_data = torch.load('intersimple-expert-data-setobs2.pt')
expert_data = Buffer(*expert_data) 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( value, policy = gail_ppo(
env_fn=env_fn, env_fn=env_fn,
expert_data=expert_data, expert_data=expert_data,
@@ -76,6 +81,7 @@ value, policy = gail_ppo(
pi_opt=pi_opt, pi_opt=pi_opt,
pi_iters=100, pi_iters=100,
logger=SummaryWriter(comment='gail-ppo-options-setobs2'), logger=SummaryWriter(comment='gail-ppo-options-setobs2'),
callback=callback,
) )
torch.save(policy.state_dict(), 'gail-ppo-options-setobs2.pt') torch.save(policy.state_dict(), 'gail-ppo-options-setobs2.pt')

View File

@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
from intersim.expert import NormalizedIntersimpleExpert from intersim.expert import NormalizedIntersimpleExpert
from wrappers import CollisionPenaltyWrapper, Minobs from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
from intersim.expert import NormalizedIntersimpleExpert from intersim.expert import NormalizedIntersimpleExpert
from wrappers import CollisionPenaltyWrapper, Minobs from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
from intersim.expert import NormalizedIntersimpleExpert from intersim.expert import NormalizedIntersimpleExpert
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
from intersim.expert import NormalizedIntersimpleExpert from intersim.expert import NormalizedIntersimpleExpert
from wrappers import CollisionPenaltyWrapper, Setobs from util.wrappers import CollisionPenaltyWrapper, Setobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
from intersim.expert import NormalizedIntersimpleExpert from intersim.expert import NormalizedIntersimpleExpert
from wrappers import CollisionPenaltyWrapper, Setobs from util.wrappers import CollisionPenaltyWrapper, Setobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
from intersim.expert import NormalizedIntersimpleExpert from intersim.expert import NormalizedIntersimpleExpert
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
env = CollisionPenaltyWrapper(IntersimpleLidarFlat( env = CollisionPenaltyWrapper(IntersimpleLidarFlat(
n_rays=5, n_rays=5,

View File

@@ -9,7 +9,7 @@ import torch.optim
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from wrappers import Minobs from util.wrappers import Minobs
obs_min = np.array([ obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.], [-1000, -1000, 0, -np.pi, -1e-1, 0.],

View File

@@ -9,7 +9,7 @@ import torch.optim
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from wrappers import Minobs from util.wrappers import Minobs
obs_min = np.array([ obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.], [-1000, -1000, 0, -np.pi, -1e-1, 0.],

View File

@@ -9,9 +9,9 @@ from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
import numpy as np 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 from options.options import OptionsEnv
obs_min = np.array([ obs_min = np.array([

View File

@@ -0,0 +1,3 @@
torch
stable-baselines3
gym

View File

@@ -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()
# %%

View File

@@ -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()
# %%

View File

@@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward
import functools import functools
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy
from wrappers import Minobs from util.wrappers import Minobs
obs_min = np.array([ obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.], [-1000, -1000, 0, -np.pi, -1e-1, 0.],

View File

@@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward
import functools import functools
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy
from wrappers import Minobs from util.wrappers import Minobs
obs_min = np.array([ obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.], [-1000, -1000, 0, -np.pi, -1e-1, 0.],

View File

@@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward
import functools import functools
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy
from wrappers import Setobs from util.wrappers import Setobs
obs_min = np.array([ obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.], [-1000, -1000, 0, -np.pi, -1e-1, 0.],

View File

@@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward
import functools import functools
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy
from wrappers import Setobs from util.wrappers import Setobs
obs_min = np.array([ obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.], [-1000, -1000, 0, -np.pi, -1e-1, 0.],

View File

@@ -9,10 +9,10 @@ from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
import numpy as np import numpy as np
from wrappers import CollisionPenaltyWrapper, TransformObservation from util.wrappers import CollisionPenaltyWrapper, TransformObservation
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy
from wrappers import Minobs from util.wrappers import Minobs
from options.options import OptionsEnv from options.options import OptionsEnv
obs_min = np.array([ obs_min = np.array([

View File

@@ -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
}

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Minobs from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Minobs from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Setobs from util.wrappers import CollisionPenaltyWrapper, Setobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat( envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat(
n_rays=5, n_rays=5,

View File

@@ -9,7 +9,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter

View File

@@ -9,7 +9,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter
@@ -50,7 +50,7 @@ value = SetValue()
v_opt = torch.optim.Adam(value.parameters(), lr=1e-4) v_opt = torch.optim.Adam(value.parameters(), lr=1e-4)
discriminator = DeepsetDiscriminator() 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 = torch.load('intersimple-expert-data-setobs2.pt')
expert_data = Buffer(*expert_data) expert_data = Buffer(*expert_data)

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Minobs from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Setobs from util.wrappers import CollisionPenaltyWrapper, Setobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter

View File

@@ -1,4 +1,3 @@
# %%
import gym import gym
from options.options import gail_ppo, Buffer from options.options import gail_ppo, Buffer
from core.value import SetValue from core.value import SetValue
@@ -8,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter
@@ -51,12 +50,11 @@ value = SetValue()
v_opt = torch.optim.Adam(value.parameters(), lr=1e-3) v_opt = torch.optim.Adam(value.parameters(), lr=1e-3)
discriminator = DeepsetDiscriminator() 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 = torch.load('intersimple-expert-data-setobs2.pt')
expert_data = Buffer(*expert_data) expert_data = Buffer(*expert_data)
# %%
value, policy = gail_ppo( value, policy = gail_ppo(
env_fn=env_fn, env_fn=env_fn,
expert_data=expert_data, expert_data=expert_data,

View File

@@ -7,7 +7,7 @@ from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
import torch import torch
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
model = PPO.load('sb3-ppo-intersimple') model = PPO.load('sb3-ppo-intersimple')
env = CollisionPenaltyWrapper(IntersimpleLidarFlat( env = CollisionPenaltyWrapper(IntersimpleLidarFlat(

View File

@@ -160,3 +160,6 @@ class ReparamPolicy(ReparamModule):
def predict(self, *args, **kwargs): def predict(self, *args, **kwargs):
return self.module.predict(*args, **kwargs) return self.module.predict(*args, **kwargs)
def unsafe_probability_mass(self, *args, **kwargs):
return self.module.unsafe_probability_mass(*args, **kwargs)

111
src/options/envs.py Normal file
View File

@@ -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

View File

@@ -18,7 +18,7 @@ class OptionsRollout:
def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value,
v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, 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(torch.zeros(env_fn(0).observation_space.shape))
policy = ReparamPolicy(policy) 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) 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) expert_data = roll_buffer(expert_data, shifts=-3, dims=0)
if callback is not None:
callback(epoch, value, policy)
return value, policy return value, policy
def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value,
v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, 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_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]) 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) 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) expert_data = roll_buffer(expert_data, shifts=-3, dims=0)
if callback is not None:
callback(epoch, value, policy)
return value, policy return value, policy
def rollout(env_fn, policy, n_episodes, max_steps_per_episode): 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)))) env = VecEnv(list(map(lambda i: (lambda: env_fn(i)), range(n_episodes))))
states[:, 0] = torch.tensor(env.reset()).clone().detach() states[:, 0] = torch.tensor(env.reset()).clone().detach()
dones[:, 0] = False
for s in tqdm(range(max_steps_per_episode), 'Rollout'): for s in tqdm(range(max_steps_per_episode), 'Rollout'):
actions[:, s] = policy.sample(policy(states[:, s])).clone().detach() 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) o, r, d, info = env.step(clipped_actions)
states[:, s + 1] = torch.tensor(o).clone().detach() states[:, s + 1] = torch.tensor(o).clone().detach()
rewards[:, s] = torch.tensor(r).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_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_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) o, r, d, i = super().step(u)
actions[k] = u actions[k] = u
rewards[k] = r rewards[k] = r
env_done[k+1] = d env_done[k] = d
infos.append(i) infos.append(i)
observations[k+1] = o 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) 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_obs = ll_obs[ll_steps]
hl_reward = (ll_rewards * ~ll_plan_done).sum().item() 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 = { hl_infos = {
'll': { 'll': {
'observations': ll_obs, 'observations': ll_obs,

View File

@@ -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)

305
src/safe_options/options.py Normal file
View File

@@ -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

View File

@@ -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)

View File

@@ -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

View File

@@ -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