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
from core.policy import SetPolicy
from wrappers import Setobs, TransformObservation, CollisionPenaltyWrapper
from util.wrappers import Setobs, TransformObservation, CollisionPenaltyWrapper
from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward
import functools

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -8,7 +8,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward
import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np
from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter
@@ -57,6 +57,11 @@ expert_data = torch.load('intersimple-expert-data-setobs2.pt')
expert_data = Buffer(*expert_data)
# %%
def callback(epoch, value, policy):
if not epoch % 10:
torch.save(policy.state_dict(), f'gail-ppo-options-setobs2-{epoch}.pt')
torch.save(value.state_dict(), f'gail-ppo-options-setobs2-value-{epoch}.pt')
value, policy = gail_ppo(
env_fn=env_fn,
expert_data=expert_data,
@@ -76,6 +81,7 @@ value, policy = gail_ppo(
pi_opt=pi_opt,
pi_iters=100,
logger=SummaryWriter(comment='gail-ppo-options-setobs2'),
callback=callback,
)
torch.save(policy.state_dict(), 'gail-ppo-options-setobs2.pt')

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

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 numpy as np
from gym.wrappers import TransformObservation
from wrappers import CollisionPenaltyWrapper
from util.wrappers import CollisionPenaltyWrapper
from core.reparam_module import ReparamPolicy
from wrappers import Minobs
from util.wrappers import Minobs
obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.],

View File

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

View File

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

View File

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

View File

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

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.intersimple import speed_reward
import functools
from wrappers import CollisionPenaltyWrapper, Minobs
from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np
from gym.wrappers import TransformObservation

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -1,201 +0,0 @@
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 Buffer, train_discriminator, roll_buffer, TerminalLogger
from dataclasses import dataclass
from core.trpo import trpo_step
from core.ppo import ppo_step
import torch.nn.functional as F
@dataclass
class OptionsRollout:
hl: Buffer
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()):
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)):
hl_data, ll_data = rollout(env_fn, policy, rollout_episodes, rollout_steps)
generator_data = OptionsRollout(Buffer(*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)
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.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)
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):
hl_data, ll_data = rollout(env_fn, policy, rollout_episodes, rollout_steps)
generator_data = OptionsRollout(Buffer(*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)
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.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)
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.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.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))))
states[:, 0] = torch.tensor(env.reset()).clone().detach()
dones[:, 0] = False
for s in tqdm(range(max_steps_per_episode), 'Rollout'):
actions[:, s] = policy.sample(policy(states[:, s])).clone().detach()
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).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]
actions = actions[:, :max_steps_per_episode]
rewards = rewards[:, :max_steps_per_episode]
dones = dones[:, :max_steps_per_episode]
return (states, 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 = []
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+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

View File

@@ -1,54 +0,0 @@
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

View File

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

View File

@@ -1,74 +0,0 @@
import numpy as np
import gym
class Wrapper(gym.Wrapper):
def __getattr__(self, name):
return getattr(self.env, name)
class TransformObservation(gym.wrappers.TransformObservation):
def __getattr__(self, name):
return getattr(self.env, name)
class CollisionPenaltyWrapper(Wrapper):
def __init__(self, env, collision_distance, collision_penalty, *args, **kwargs):
super().__init__(env, *args, **kwargs)
self.penalty = collision_penalty
self.distance = collision_distance
def step(self, action):
obs, reward, done, info = super().step(action)
reward = -self.penalty if (obs.reshape(-1, 6)[1:, 0] < self.distance).any() else reward
self.env._rewards.pop()
self.env._rewards.append(reward)
return obs, reward, done, info
class Minobs(Wrapper):
""" Meant to be used as wrapper around LidarObservation """
def __init__(self, env, *args, **kwargs):
super().__init__(env, *args, **kwargs)
n_rays = int(self.observation_space.shape[0] / 6) - 1
self.observation_space = gym.spaces.Box(low=-np.inf, high=np.inf, shape=((1 + n_rays) * 2,))
def minobs(self, obs):
""" ego v, psidot ; (for each ray,) rel. distance, rel. velocity in ego forward direction """
obs = obs.reshape(-1, 6)
obs = np.concatenate((obs[:1, [2, 4]], obs[1:, [0, 2]]), axis=0)
return obs.reshape(-1)
def reset(self):
return self.minobs(super().reset())
def step(self, action):
obs, reward, done, info = super().step(action)
return self.minobs(obs), reward, done, info
class Setobs(Wrapper):
""" Meant to be used as wrapper around LidarObservation """
def __init__(self, env, *args, **kwargs):
super().__init__(env, *args, **kwargs)
self.n_rays = int(self.observation_space.shape[0] / 6) - 1
self.observation_space = gym.spaces.Box(low=-np.inf, high=np.inf, shape=(self.n_rays, 6))
def obs(self, obs):
obs = obs.reshape(-1, 6)
ego = obs[:1, [2, 4]] # v, psidot
ego = np.tile(ego, (self.n_rays, 1))
other = obs[1:, [0, 1, 2]] # distance, angle, velocity component in ego forward direction
other = np.stack((other[:, 0], np.cos(other[:, 1]), np.sin(other[:, 1]), other[:, 2]), axis=-1)
obs = np.concatenate((ego, other), axis=-1)
return obs
def reset(self):
return self.obs(super().reset())
def step(self, action):
obs, reward, done, info = super().step(action)
return self.obs(obs), reward, done, info