From 072c0ff417e28b6287369f916cb92e9ee02dc9d7 Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Tue, 15 Feb 2022 14:03:22 +0100 Subject: [PATCH] Copy files --- src/core/discriminator.py | 74 ++++++++++++++++ src/core/gail.py | 125 ++++++++++++++++++++++++++ src/core/optimization.py | 39 ++++++++ src/core/policy.py | 96 ++++++++++++++++++++ src/core/ppo.py | 72 +++++++++++++++ src/core/reparam_module.py | 162 ++++++++++++++++++++++++++++++++++ src/core/sampling.py | 73 +++++++++++++++ src/core/test_optimization.py | 23 +++++ src/core/trpo.py | 79 +++++++++++++++++ src/core/value.py | 48 ++++++++++ src/core/value_estimation.py | 40 +++++++++ src/gail2/envs.py | 104 ++++++++++++++++++++++ src/gail2/options.py | 128 +++++++++++++++++++++++++++ src/gail2/test_options.py | 54 ++++++++++++ src/gail2/wrappers.py | 74 ++++++++++++++++ 15 files changed, 1191 insertions(+) create mode 100644 src/core/discriminator.py create mode 100644 src/core/gail.py create mode 100644 src/core/optimization.py create mode 100644 src/core/policy.py create mode 100644 src/core/ppo.py create mode 100644 src/core/reparam_module.py create mode 100644 src/core/sampling.py create mode 100644 src/core/test_optimization.py create mode 100644 src/core/trpo.py create mode 100644 src/core/value.py create mode 100644 src/core/value_estimation.py create mode 100644 src/gail2/envs.py create mode 100644 src/gail2/options.py create mode 100644 src/gail2/test_options.py create mode 100644 src/gail2/wrappers.py diff --git a/src/core/discriminator.py b/src/core/discriminator.py new file mode 100644 index 0000000..8074c32 --- /dev/null +++ b/src/core/discriminator.py @@ -0,0 +1,74 @@ +import torch +import torch.nn as nn + +class Discriminator(nn.Module): + + def __init__(self): + super().__init__() + self.nn = nn.Sequential( + nn.LazyLinear(50), + nn.Tanh(), + nn.LazyLinear(50), + nn.Tanh(), + nn.LazyLinear(1), + ) + + def forward(self, states, actions): + return self.nn(torch.cat((states, actions), dim=-1)).squeeze(-1) + +class DeepsetDiscriminator(nn.Module): + + def __init__(self): + super().__init__() + self.elem = nn.Sequential( + nn.LazyLinear(10), + nn.Tanh(), + nn.LazyLinear(10), + nn.Tanh(), + nn.LazyLinear(10), + ) + self.glob = nn.Sequential( + nn.LazyLinear(10), + nn.Tanh(), + nn.LazyLinear(10), + nn.Tanh(), + nn.LazyLinear(1), + ) + + def forward(self, states, actions): + actions = actions.unsqueeze(-2) + actions = actions.expand(*actions.shape[:-2], states.shape[-2], actions.shape[-1]) + sa = torch.cat((states, actions), dim=-1) + return self.glob(self.elem(sa).sum(-2)).squeeze(-1) + +class RecurrentDiscriminator(nn.Module): + + def __init__(self): + super().__init__() + self.state_dim = 10 + self.state = nn.Sequential( + nn.LazyLinear(10), + nn.Tanh(), + nn.LazyLinear(10), + nn.Tanh(), + nn.LazyLinear(self.state_dim), + ) + self.glob = nn.Sequential( + nn.LazyLinear(10), + nn.Tanh(), + nn.LazyLinear(1), + ) + + def forward(self, states, actions): + actions = actions.unsqueeze(-2) + batch_size = actions.shape[:-2] + set_size = states.shape[-2] + action_dim = actions.shape[-1] + actions = actions.expand(*batch_size, set_size, action_dim) + sa = torch.cat((states, actions), dim=-1) + + state = torch.zeros((*batch_size, self.state_dim)) + for i in range(set_size): + state = state + self.state(torch.cat((state, sa[..., i, :]), dim=-1)) + + return self.glob(state).squeeze(-1) diff --git a/src/core/gail.py b/src/core/gail.py new file mode 100644 index 0000000..d630bd4 --- /dev/null +++ b/src/core/gail.py @@ -0,0 +1,125 @@ +import torch +import torch.nn.functional as F +from dataclasses import dataclass +from core.reparam_module import ReparamPolicy +from core.sampling import rollout +from core.trpo import trpo_step +from core.ppo import ppo_step +from tqdm import tqdm + +class TerminalLogger: + def add_scalar(self, key, scalar, i=None): + if i is not None: + print('Iteration', i, end=' ') + print(key, scalar) + +@dataclass +class Buffer: + states: torch.Tensor + actions: torch.Tensor + rewards: torch.Tensor + dones: torch.Tensor + +def roll_buffer(buffer, *args, **kwargs): + return Buffer( + torch.roll(buffer.states, *args, **kwargs), + torch.roll(buffer.actions, *args, **kwargs), + torch.roll(buffer.rewards, *args, **kwargs), + torch.roll(buffer.dones, *args, **kwargs), + ) + +def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, + v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, + gae_lambda, delta, backtrack_coeff, backtrack_iters, cg_iters=10, cg_damping=0.1, wasserstein=False, wasserstein_c=None, logger=TerminalLogger()): + + policy(torch.zeros(env_fn(0).observation_space.shape)) + policy = ReparamPolicy(policy) + + logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0]) + logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0]) + + for epoch in tqdm(range(epochs)): + generator_data = Buffer(*rollout(env_fn, policy, rollout_episodes, rollout_steps)) + + logger.add_scalar('gen/mean_episode_length', (~generator_data.dones).sum() / generator_data.states.shape[0], epoch) + logger.add_scalar('gen/mean_reward_per_episode', generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0], epoch) + + discriminator, loss = train_discriminator(expert_data, generator_data, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c) + if wasserstein: + generator_data.rewards = discriminator(generator_data.states, generator_data.actions) + else: + generator_data.rewards = -F.logsigmoid(discriminator(generator_data.states, generator_data.actions)) + logger.add_scalar('disc/final_loss', loss, epoch) + logger.add_scalar('disc/mean_reward_per_episode', generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0], epoch) + + value, policy = trpo_step(value, policy, generator_data.states, generator_data.actions, generator_data.rewards, generator_data.dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters, cg_damping) + expert_data = roll_buffer(expert_data, shifts=-3, dims=0) + + return value, policy + +def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, + v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, + gae_lambda, clip_ratio, pi_opt, pi_iters, target_kl=None, max_grad_norm=None, wasserstein=False, wasserstein_c=None, logger=TerminalLogger()): + + logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0]) + logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0]) + + for epoch in range(epochs): + generator_data = Buffer(*rollout(env_fn, policy, rollout_episodes, rollout_steps)) + + logger.add_scalar('gen/mean_episode_length', (~generator_data.dones).sum() / generator_data.states.shape[0], epoch) + logger.add_scalar('gen/mean_reward_per_episode', generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0], epoch) + + discriminator, loss = train_discriminator(expert_data, generator_data, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c) + if wasserstein: + generator_data.rewards = discriminator(generator_data.states, generator_data.actions) + else: + generator_data.rewards = -F.logsigmoid(discriminator(generator_data.states, generator_data.actions)) + logger.add_scalar('disc/final_loss', loss, epoch) + logger.add_scalar('disc/mean_reward_per_episode', generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0], epoch) + + value, policy = ppo_step(value, policy, generator_data.states, generator_data.actions, generator_data.rewards, generator_data.dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm) + expert_data = roll_buffer(expert_data, shifts=-3, dims=0) + + return value, policy + +def train_discriminator(expert_data, generator_data, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c=None): + + n_expert_samples = (~expert_data.dones).sum() + n_generator_samples = (~generator_data.dones).sum() + n_samples = torch.minimum(n_expert_samples, n_generator_samples) + + gen_states = generator_data.states[~generator_data.dones][:n_samples] + gen_actions = generator_data.actions[~generator_data.dones][:n_samples] + exp_states = expert_data.states[~expert_data.dones][:n_samples] + exp_actions = expert_data.actions[~expert_data.dones][:n_samples] + + states = torch.cat((exp_states, gen_states), dim=0).detach() + actions = torch.cat((exp_actions, gen_actions), dim=0).detach() + labels = torch.cat((torch.zeros(n_samples), torch.ones(n_samples))).detach() + + # print('Batch augmentation on') + # random_states = torch.rand_like(gen_states) + # random_actions = torch.rand_like(gen_actions) + # states = torch.cat((exp_states, gen_states, random_states), dim=0).detach() + # actions = torch.cat((exp_actions, gen_actions, random_actions), dim=0).detach() + # labels = torch.cat((torch.zeros(n_samples), torch.ones(n_samples), torch.ones(n_samples))).detach() + + for _ in range(disc_iters): + disc_opt.zero_grad() + pred = discriminator(states, actions) + + if wasserstein: + loss = -(pred * (1 - labels) - pred * labels).mean() + else: + loss = F.binary_cross_entropy(torch.sigmoid(pred), labels) + + loss.backward() + disc_opt.step() + + if wasserstein_c is not None: + with torch.no_grad(): + for param in discriminator.parameters(): + param.clamp_(-wasserstein_c, wasserstein_c) + + return discriminator, loss diff --git a/src/core/optimization.py b/src/core/optimization.py new file mode 100644 index 0000000..4057214 --- /dev/null +++ b/src/core/optimization.py @@ -0,0 +1,39 @@ +import torch + +def conjugate_gradient(A, b, max_iters, res_tol=1e-10): + x = torch.zeros_like(b) + r = b - A(x) + p = r + + rTr = r.T @ r + + for _ in range(max_iters): + Ap = A(p) + alpha = rTr / (p.T @ Ap) + x = x + alpha * p + + r = r - alpha * Ap + if torch.norm(r) < res_tol: + break + + rTrnew = r.T @ r + beta = rTrnew / rTr + p = r + beta * p + rTr = rTrnew + + return x + +def line_search(f, x0, dx, g0, alpha, condition, max_steps=10, c1=0.1): + assert 0 < alpha < 1 + + f0 = f(x0) + for _ in range(max_steps): + x = x0 + dx + + if (f(x) > f0 + c1 * g0.T @ dx) and condition(x): + return x + + dx *= alpha + + print('Line search failed, returning x0') + return x0 diff --git a/src/core/policy.py b/src/core/policy.py new file mode 100644 index 0000000..579e72b --- /dev/null +++ b/src/core/policy.py @@ -0,0 +1,96 @@ +import torch +import torch.nn as nn +from torch.distributions import Independent, Normal, Categorical +from torch.distributions.kl import kl_divergence + +class BasePolicy(nn.Module): + + def __init__(self, action_dim): + super().__init__() + self.action_dim = action_dim + + def torch_dist(self, dist): + return Independent(Normal(dist[..., :self.action_dim], dist[..., self.action_dim:].exp()), 1) + + def sample(self, dist): + return self.torch_dist(dist).sample() + + def predict(self, states): + return self.sample(self.forward(states)) + + def log_prob(self, dist, actions): + return self.torch_dist(dist).log_prob(actions) + + def kl_divergence(self, dist1, dist2): + d1 = self.torch_dist(dist1) + d2 = self.torch_dist(dist2) + return kl_divergence(d1, d2) + +class Policy(BasePolicy): + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.nn = nn.Sequential( + nn.LazyLinear(50), + nn.Tanh(), + nn.LazyLinear(50), + nn.Tanh(), + nn.LazyLinear(2 * self.action_dim), + ) + + def forward(self, states): + return self.nn(states) + +class DiscretePolicy(BasePolicy): + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.nn = nn.Sequential( + nn.LazyLinear(50), + nn.Tanh(), + nn.LazyLinear(50), + nn.Tanh(), + nn.LazyLinear(self.action_dim), + ) + + def forward(self, states): + return self.nn(states) + + def torch_dist(self, dist): + return Categorical(logits=dist) + +class SetPolicy(Policy): + + def forward(self, states): + batch_size = states.shape[:-2] + states = torch.cat((states[..., :1, [0, 1]], states[..., :, [2, 5]]), axis=-2).reshape(*batch_size, -1) + return super().forward(states) + +class SetDiscretePolicy(DiscretePolicy): + + def forward(self, states): + batch_size = states.shape[:-2] + states = torch.cat((states[..., :1, [0, 1]], states[..., :, [2, 5]]), axis=-2).reshape(*batch_size, -1) + return super().forward(states) + +class DeepSetPolicy(BasePolicy): + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.elem = nn.Sequential( + nn.LazyLinear(10), + nn.Tanh(), + nn.LazyLinear(10), + nn.Tanh(), + nn.LazyLinear(10), + ) + self.glob = nn.Sequential( + nn.LazyLinear(10), + nn.Tanh(), + nn.LazyLinear(10), + nn.Tanh(), + nn.LazyLinear(2 * self.action_dim), + ) + + def forward(self, states): + return self.glob(self.elem(states).sum(-2)) diff --git a/src/core/ppo.py b/src/core/ppo.py new file mode 100644 index 0000000..c742061 --- /dev/null +++ b/src/core/ppo.py @@ -0,0 +1,72 @@ +import torch +from core.sampling import rollout +from core.value_estimation import gae + +def ppo(env_fn, value, policy, epochs, rollout_episodes, rollout_steps, gamma, gae_lambda, clip_ratio, pi_opt, pi_iters, v_opt, v_iters, target_kl=None, max_grad_norm=None): + + for epoch in range(epochs): + policy.eval() + states, actions, rewards, dones = rollout(env_fn, policy, rollout_episodes, rollout_steps) + + print('mean', states[~dones].mean(0)) + print('std', states[~dones].std(0)) + + print(f'Iteration {epoch} mean episode length {(~dones).sum() / states.shape[0]}') + print(f'Iteration {epoch} mean reward per episode {rewards[~dones].sum() / states.shape[0]}') + + policy.train() + value.train() + value, policy = ppo_step(value, policy, states, actions, rewards, dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm) + + return value, policy + +def ppo_step(value, policy, states, actions, rewards, dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm): + + states = states.detach() + actions = actions.detach() + rewards = rewards.detach() + dones = dones.detach() + + advantages, returns, valid = gae(states, rewards, value(states), dones, gamma, gae_lambda) + advantages = advantages.detach() + returns = returns.detach() + + # update value function + + for _ in range(v_iters): + v_opt.zero_grad() + value_loss = (value(states) - returns).pow(2)[valid].mean() + value_loss.backward() + v_opt.step() + + # update policy + + old_dist = policy(states).detach() + old_logprob = policy.log_prob(old_dist, actions).detach() + + def g(advantages, clip_ratio): + return torch.where(advantages >= 0, (1 + clip_ratio) * advantages, (1 - clip_ratio) * advantages) + + def L(states, actions, advantages, clip_ratio): + return torch.minimum( + (policy.log_prob(policy(states), actions) - old_logprob).exp() * advantages, + g(advantages, clip_ratio) + )[valid].mean() + + for _ in range(pi_iters): + pi_opt.zero_grad() + ppo_loss = -L(states, actions, advantages, clip_ratio) + ppo_loss.backward() + + if max_grad_norm: + torch.nn.utils.clip_grad_norm(policy.parameters(), max_grad_norm) + + pi_opt.step() + + kl = policy.kl_divergence(policy(states), old_dist)[valid].mean() + if target_kl and kl > target_kl: + break + + print('KL', kl.item()) + + return value, policy diff --git a/src/core/reparam_module.py b/src/core/reparam_module.py new file mode 100644 index 0000000..5bcd613 --- /dev/null +++ b/src/core/reparam_module.py @@ -0,0 +1,162 @@ +# Source: https://github.com/SsnL/PyTorch-Reparam-Module + +import torch +import torch.nn as nn +import warnings +import types +from collections import namedtuple +from contextlib import contextmanager + +class ReparamModule(nn.Module): + def __init__(self, module): + super(ReparamModule, self).__init__() + self.module = module + + param_infos = [] + shared_param_memo = {} + shared_param_infos = [] + params = [] + param_numels = [] + param_shapes = [] + for m in self.modules(): + for n, p in m.named_parameters(recurse=False): + if p is not None: + if p in shared_param_memo: + shared_m, shared_n = shared_param_memo[p] + shared_param_infos.append((m, n, shared_m, shared_n)) + else: + shared_param_memo[p] = (m, n) + param_infos.append((m, n)) + params.append(p.detach()) + param_numels.append(p.numel()) + param_shapes.append(p.size()) + + assert len(set(p.dtype for p in params)) <= 1, \ + "expects all parameters in module to have same dtype" + + # store the info for unflatten + self._param_infos = tuple(param_infos) + self._shared_param_infos = tuple(shared_param_infos) + self._param_numels = tuple(param_numels) + self._param_shapes = tuple(param_shapes) + + # flatten + flat_param = nn.Parameter(torch.cat([p.reshape(-1) for p in params], 0)) + self.register_parameter('flat_param', flat_param) + self.param_numel = flat_param.numel() + del params + del shared_param_memo + + # deregister the names as parameters + for m, n in self._param_infos: + delattr(m, n) + for m, n, _, _ in self._shared_param_infos: + delattr(m, n) + + # register the views as plain attributes + self._unflatten_param(self.flat_param) + + # now buffers + # they are not reparametrized. just store info as (module, name, buffer) + buffer_infos = [] + for m in self.modules(): + for n, b in m.named_buffers(recurse=False): + if b is not None: + buffer_infos.append((m, n, b)) + + self._buffer_infos = tuple(buffer_infos) + self._traced_self = None + + def trace(self, example_input, **trace_kwargs): + assert self._traced_self is None, 'This ReparamModule is already traced' + + if isinstance(example_input, torch.Tensor): + example_input = (example_input,) + example_input = tuple(example_input) + example_param = (self.flat_param.detach().clone(),) + example_buffers = (tuple(b.detach().clone() for _, _, b in self._buffer_infos),) + + self._traced_self = torch.jit.trace_module( + self, + inputs=dict( + _forward_with_param=example_param + example_input, + _forward_with_param_and_buffers=example_param + example_buffers + example_input, + ), + **trace_kwargs, + ) + + # replace forwards with traced versions + self._forward_with_param = self._traced_self._forward_with_param + self._forward_with_param_and_buffers = self._traced_self._forward_with_param_and_buffers + return self + + def clear_views(self): + for m, n in self._param_infos: + setattr(m, n, None) # This will set as plain attr + + def _apply(self, *args, **kwargs): + if self._traced_self is not None: + self._traced_self._apply(*args, **kwargs) + return self + return super(ReparamModule, self)._apply(*args, **kwargs) + + def _unflatten_param(self, flat_param): + ps = (t.view(s) for (t, s) in zip(flat_param.split(self._param_numels), self._param_shapes)) + for (m, n), p in zip(self._param_infos, ps): + setattr(m, n, p) # This will set as plain attr + for (m, n, shared_m, shared_n) in self._shared_param_infos: + setattr(m, n, getattr(shared_m, shared_n)) + + @contextmanager + def unflattened_param(self, flat_param): + saved_views = [getattr(m, n) for m, n in self._param_infos] + self._unflatten_param(flat_param) + yield + # Why not just `self._unflatten_param(self.flat_param)`? + # 1. because of https://github.com/pytorch/pytorch/issues/17583 + # 2. slightly faster since it does not require reconstruct the split+view + # graph + for (m, n), p in zip(self._param_infos, saved_views): + setattr(m, n, p) + for (m, n, shared_m, shared_n) in self._shared_param_infos: + setattr(m, n, getattr(shared_m, shared_n)) + + @contextmanager + def replaced_buffers(self, buffers): + for (m, n, _), new_b in zip(self._buffer_infos, buffers): + setattr(m, n, new_b) + yield + for m, n, old_b in self._buffer_infos: + setattr(m, n, old_b) + + def _forward_with_param_and_buffers(self, flat_param, buffers, *inputs, **kwinputs): + with self.unflattened_param(flat_param): + with self.replaced_buffers(buffers): + return self.module(*inputs, **kwinputs) + + def _forward_with_param(self, flat_param, *inputs, **kwinputs): + with self.unflattened_param(flat_param): + return self.module(*inputs, **kwinputs) + + def forward(self, *inputs, flat_param=None, buffers=None, **kwinputs): + if flat_param is None: + flat_param = self.flat_param + if buffers is None: + return self._forward_with_param(flat_param, *inputs, **kwinputs) + else: + return self._forward_with_param_and_buffers(flat_param, tuple(buffers), *inputs, **kwinputs) + + +class ReparamPolicy(ReparamModule): + + def sample(self, *args, **kwargs): + return self.module.sample(*args, **kwargs) + + def log_prob(self, *args, **kwargs): + return self.module.log_prob(*args, **kwargs) + + def kl_divergence(self, *args, **kwargs): + return self.module.kl_divergence(*args, **kwargs) + + def predict(self, *args, **kwargs): + return self.module.predict(*args, **kwargs) diff --git a/src/core/sampling.py b/src/core/sampling.py new file mode 100644 index 0000000..66fbbef --- /dev/null +++ b/src/core/sampling.py @@ -0,0 +1,73 @@ +import torch +import gym +from stable_baselines3.common.vec_env import DummyVecEnv as VecEnv +from tqdm import tqdm + +def rollout(env_fn, policy, n_episodes, max_steps_per_episode): + env = env_fn(0) + states = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.observation_space.shape) + actions = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.action_space.shape) + rewards = torch.zeros(n_episodes, max_steps_per_episode + 1) + dones = torch.ones(n_episodes, max_steps_per_episode + 1, dtype=bool) + + env = VecEnv(list(map(lambda i: (lambda: env_fn(i)), range(n_episodes)))) + + states[:, 0] = torch.tensor(env.reset()).clone().detach() + dones[:, 0] = False + + for s in range(max_steps_per_episode): + actions[:, s] = policy.sample(policy(states[:, s])).clone().detach() + + clipped_actions = actions[:, s] + if isinstance(env.action_space, gym.spaces.Box): + clipped_actions = torch.clamp(clipped_actions, torch.from_numpy(env.action_space.low), torch.from_numpy(env.action_space.high)) + + o, r, d, _ = env.step(clipped_actions) + states[:, s + 1] = torch.tensor(o).clone().detach() + rewards[:, s] = torch.tensor(r).clone().detach() + dones[:, s + 1] = torch.tensor(d).clone().detach() + + dones = dones.cumsum(1) > 0 + + states = states[:, :max_steps_per_episode] + actions = actions[:, :max_steps_per_episode] + rewards = rewards[:, :max_steps_per_episode] + dones = dones[:, :max_steps_per_episode] + + return states, actions, rewards, dones + + +def rollout_sb3(env, policy, n_episodes, max_steps_per_episode): + states = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.observation_space.shape) + actions = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.action_space.shape) + rewards = torch.zeros(n_episodes, max_steps_per_episode + 1) + dones = torch.ones(n_episodes, max_steps_per_episode + 1, dtype=bool) + + for e in tqdm(range(n_episodes)): + states[e, 0] = torch.tensor(env.reset()).clone().detach() + dones[e, 0] = False + + for s in range(max_steps_per_episode): + action, _ = policy.predict(states[e, s]) + actions[e, s] = torch.tensor(action).clone().detach() + + clipped_actions = actions[e, s] + if isinstance(env.action_space, gym.spaces.Box): + clipped_actions = torch.clamp(clipped_actions, torch.from_numpy(env.action_space.low), torch.from_numpy(env.action_space.high)) + + o, r, d, _ = env.step(clipped_actions) + states[e, s + 1] = torch.tensor(o).clone().detach() + rewards[e, s] = torch.tensor(r).clone().detach() + dones[e, s + 1] = torch.tensor(d).clone().detach() + + if d: + break + + dones = dones.cumsum(1) > 0 + + states = states[:, :max_steps_per_episode] + actions = actions[:, :max_steps_per_episode] + rewards = rewards[:, :max_steps_per_episode] + dones = dones[:, :max_steps_per_episode] + + return states, actions, rewards, dones diff --git a/src/core/test_optimization.py b/src/core/test_optimization.py new file mode 100644 index 0000000..aa49bb5 --- /dev/null +++ b/src/core/test_optimization.py @@ -0,0 +1,23 @@ +import torch +from optimization import conjugate_gradient + +def test_cg_eye(): + A = torch.eye(2) + b = torch.tensor([1., 2.]) + x1 = conjugate_gradient(lambda x: A @ x, b, 2) + x2 = torch.inverse(A) @ b + assert torch.allclose(x1, x2) + +def test_cg_eyep1(): + A = torch.eye(2) + 1 + b = torch.tensor([1., 2.]) + x1 = conjugate_gradient(lambda x: A @ x, b, 2) + x2 = torch.inverse(A) @ b + assert torch.allclose(x1, x2, atol=1e-7) + +def test_cg3(): + A = torch.tensor([[4., 2.], [2., 4.]]) + b = torch.tensor([2., 1.]) + x1 = conjugate_gradient(lambda x: A @ x, b, 100) + x2 = torch.inverse(A) @ b + assert torch.allclose(x1, x2) diff --git a/src/core/trpo.py b/src/core/trpo.py new file mode 100644 index 0000000..d91e363 --- /dev/null +++ b/src/core/trpo.py @@ -0,0 +1,79 @@ +import torch +from core.reparam_module import ReparamPolicy +from core.sampling import rollout +from core.value_estimation import gae +from core.optimization import conjugate_gradient, line_search + +def trpo(env_fn, value, policy, epochs, rollout_episodes, rollout_steps, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1): + + policy(torch.zeros(env_fn(0).observation_space.shape)) + policy = ReparamPolicy(policy) + + for epoch in range(epochs): + policy.eval() + states, actions, rewards, dones = rollout(env_fn, policy, rollout_episodes, rollout_steps) + + print('mean', states[~dones].mean(0)) + print('std', states[~dones].std(0)) + + print(f'Iteration {epoch} mean episode length {(~dones).sum() / states.shape[0]}') + print(f'Iteration {epoch} mean reward per episode {rewards[~dones].sum() / states.shape[0]}') + + policy.train() + value.train() + value, policy = trpo_step(value, policy, states, actions, rewards, dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters, cg_damping) + + return value, policy + +def trpo_step(value, policy, states, actions, rewards, dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1): + + states = states.detach() + actions = actions.detach() + rewards = rewards.detach() + dones = dones.detach() + + advantages, returns, valid = gae(states, rewards, value(states), dones, gamma, gae_lambda) + advantages = advantages.detach() + returns = returns.detach() + + # update value function + + for _ in range(v_iters): + v_opt.zero_grad() + value_loss = (value(states) - returns).pow(2)[valid].mean() + value_loss.backward() + v_opt.step() + + # compute policy gradient + + plogprob = policy.log_prob(policy(states), actions) + surrogate_advantage = (plogprob * advantages)[valid].sum() / states.shape[0] + g = torch.cat(torch.autograd.grad(surrogate_advantage, policy.flat_param)).detach() + + def Hx(x): + kl = policy.kl_divergence(policy(states), policy(states).detach())[valid].mean() + dKL = torch.cat(torch.autograd.grad(kl, policy.flat_param, create_graph=True)) + H_x = torch.cat(torch.autograd.grad(dKL.T @ x, policy.flat_param)).detach() + return H_x + cg_damping * x + + x = conjugate_gradient(Hx, g, cg_iters) + npg = torch.sqrt(2 * delta / (x.T @ Hx(x))) * x + + # perform line search + + def L(theta): + rplogprob = policy.log_prob(policy(states, flat_param=theta), actions) + return ((rplogprob - plogprob.detach()).exp() * advantages)[valid].sum() / advantages.shape[0] + + condition = lambda theta: policy.kl_divergence(policy(states, flat_param=theta), policy(states))[valid].mean() < delta + + x0 = policy.flat_param + g0 = torch.cat(torch.autograd.grad(L(x0), x0)) + theta = line_search(L, x0, npg, g0, backtrack_coeff, condition, max_steps=backtrack_iters) + + # update policy parameters + + with torch.no_grad(): + policy.flat_param.copy_(theta) + + return value, policy diff --git a/src/core/value.py b/src/core/value.py new file mode 100644 index 0000000..2be3b78 --- /dev/null +++ b/src/core/value.py @@ -0,0 +1,48 @@ +import torch +import torch.nn as nn +from torch.distributions import Normal +from torch.distributions.kl import kl_divergence + +class Value(nn.Module): + + def __init__(self): + super().__init__() + self.nn = nn.Sequential( + nn.LazyLinear(50), + nn.Tanh(), + nn.LazyLinear(50), + nn.Tanh(), + nn.LazyLinear(1), + ) + + def forward(self, states): + return self.nn(states).squeeze(-1) + +class SetValue(Value): + + def forward(self, states): + batch_size = states.shape[:-2] + states = torch.cat((states[..., :1, [0, 1]], states[..., :, [2, 5]]), axis=-2).reshape(*batch_size, -1) + return super().forward(states) + +class DeepSetValue(nn.Module): + + def __init__(self): + super().__init__() + self.elem = nn.Sequential( + nn.LazyLinear(10), + nn.Tanh(), + nn.LazyLinear(10), + nn.Tanh(), + nn.LazyLinear(10), + ) + self.glob = nn.Sequential( + nn.LazyLinear(10), + nn.Tanh(), + nn.LazyLinear(10), + nn.Tanh(), + nn.LazyLinear(1), + ) + + def forward(self, states): + return self.glob(self.elem(states).sum(-2)).squeeze(-1) diff --git a/src/core/value_estimation.py b/src/core/value_estimation.py new file mode 100644 index 0000000..b558ac2 --- /dev/null +++ b/src/core/value_estimation.py @@ -0,0 +1,40 @@ +from operator import index +import torch + +def gae(states, rewards, values, dones, gamma, gae_lambda): + assert rewards.shape == values.shape == dones.shape + n_episodes, n_steps = rewards.shape + + valid = ~dones + valid[..., -1] = False + + td = rewards + gamma * torch.roll(values, shifts=-1, dims=1) - values + adv = td.repeat(n_steps, 1, 1).transpose(0, 1) + assert adv.shape == (n_episodes, n_steps, n_steps) + + step_start, step = torch.meshgrid(torch.arange(n_steps), torch.arange(n_steps), indexing='ij') + past = step < step_start + + # add up discounted temporal differences + discount = torch.minimum(torch.tensor(gamma).log() * (step - step_start), torch.tensor(0.)).exp() + discount = discount * ~past + discount = discount * valid.unsqueeze(1) + + adv = adv * discount + adv = adv.cumsum(2) # eq. (14) + assert adv.shape == (n_episodes, n_steps, n_steps) + + # add up discounted k-advantages + lambda_discount = torch.minimum(torch.tensor(gae_lambda).log() * (step - step_start), torch.tensor(0.)).exp() + lambda_discount = lambda_discount * ~past + lambda_discount = lambda_discount * valid.unsqueeze(1) + + adv = adv * lambda_discount + adv = adv.sum(2) / (lambda_discount.sum(2) + 1e-10) # eq. (16) + + adv = (adv - adv[valid].mean()) / adv[valid].std() + assert adv.shape == rewards.shape == values.shape + + returns = adv + values + + return adv, returns, valid diff --git a/src/gail2/envs.py b/src/gail2/envs.py new file mode 100644 index 0000000..9a48f01 --- /dev/null +++ b/src/gail2/envs.py @@ -0,0 +1,104 @@ +import gym +import numpy as np +from wrappers import 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, + stop_on_collision=False, + **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)]) + +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 diff --git a/src/gail2/options.py b/src/gail2/options.py new file mode 100644 index 0000000..e2d7263 --- /dev/null +++ b/src/gail2/options.py @@ -0,0 +1,128 @@ +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) diff --git a/src/gail2/test_options.py b/src/gail2/test_options.py new file mode 100644 index 0000000..269eb43 --- /dev/null +++ b/src/gail2/test_options.py @@ -0,0 +1,54 @@ +from intersim.envs import IntersimpleLidarFlat +from options import OptionsEnv +import gym +import numpy as np + +def test_obs_shape(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + assert env.reset().shape == (36,) + +def test_act_space(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + assert env.action_space == gym.spaces.Discrete(3) + +def test_plan(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + env.reset() + plan = env.plan(options[0]) + assert np.allclose(plan, -13.998268127441406 * np.ones((5,))) + +def test_plan2(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + obs = env.reset() + states, actions, rewards, dones, plan_done, infos, n_steps = env.execute_plan(obs, options[0]) + assert states.shape == (6, 36) + assert rewards.shape == (6,) + assert dones.shape == (6,) + assert len(infos) == 5 + +def test_step(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + env.reset() + obs, reward, done, _ = env.step(0) + assert obs.shape == (36,) + assert reward == 5.0 + assert done == False + +def test_ll_step(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + env.reset() + _, _, _, info = env.step(0) + assert info['ll']['observations'].shape == (6, 36) + assert info['ll']['actions'].shape == (6, 1) + assert info['ll']['rewards'].shape == (6,) + assert info['ll']['env_done'].shape == (6,) + assert info['ll']['plan_done'].shape == (6,) + assert info['ll']['plan_done'][5] == True + assert info['ll']['steps'] == 5 + assert len(info['ll']['infos']) == 5 diff --git a/src/gail2/wrappers.py b/src/gail2/wrappers.py new file mode 100644 index 0000000..3916e75 --- /dev/null +++ b/src/gail2/wrappers.py @@ -0,0 +1,74 @@ +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