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 01/10] 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 From c6a4c10605394ec13e9f964b7694276344fc9b44 Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Tue, 15 Feb 2022 18:36:53 +0100 Subject: [PATCH 02/10] Integrate options env and policy --- evaluate_models.sh | 2 +- src/core/policy.py | 17 +++++++++++++++-- src/eval_main.py | 29 ++++++++++++++++++++++++++--- src/evaluation/evaluation.py | 17 ++++++++++++++++- src/gail2/envs.py | 4 ++-- 5 files changed, 60 insertions(+), 9 deletions(-) diff --git a/evaluate_models.sh b/evaluate_models.sh index 15ff738..f1a4b62 100755 --- a/evaluate_models.sh +++ b/evaluate_models.sh @@ -13,4 +13,4 @@ python -m src.eval_main # idm python -m src.eval_main --method=idm - +python -m src.eval_main --method=ogail --policy_file='checkpoints/gail-options-setobs2.pt' --env='NormalizedOptionsEvalEnv' diff --git a/src/core/policy.py b/src/core/policy.py index 579e72b..5e377c9 100644 --- a/src/core/policy.py +++ b/src/core/policy.py @@ -15,8 +15,13 @@ class BasePolicy(nn.Module): def sample(self, dist): return self.torch_dist(dist).sample() - def predict(self, states): - return self.sample(self.forward(states)) + def predict(self, observations, state=None, episode_start=None, deterministic=True): + observations = torch.tensor(observations) + if deterministic: + actions = self.forward(observations)[..., :self.action_dim] + else: + actions = self.sample(self.forward(observations)) + return actions, None def log_prob(self, dist, actions): return self.torch_dist(dist).log_prob(actions) @@ -58,6 +63,14 @@ class DiscretePolicy(BasePolicy): def torch_dist(self, dist): return Categorical(logits=dist) + + def predict(self, observations, state=None, episode_start=None, deterministic=True): + observations = torch.tensor(observations) + if deterministic: + _, actions = self.forward(observations).max(-1) + else: + actions = self.sample(self.forward(observations)) + return actions, None class SetPolicy(Policy): diff --git a/src/eval_main.py b/src/eval_main.py index edef877..b77bf91 100644 --- a/src/eval_main.py +++ b/src/eval_main.py @@ -8,6 +8,9 @@ from src.baselines import IDMRulePolicy from src.evaluation import IntersimpleEvaluation import src.gail.options as options_envs from src.evaluation.metrics import divergence, visualize_distribution +from src.core.policy import SetPolicy, SetDiscretePolicy +from src.core.reparam_module import ReparamPolicy +from src.gail2 import envs as options_envs2 from typing import Optional, List, Dict, Tuple import torch @@ -33,12 +36,31 @@ def load_policy(method:str, if method == 'idm': policy = IDMRulePolicy(env, **policy_kwargs) elif method == 'bc': - raise NotImplementedError + policy = SetPolicy(env.action_space.shape[-1]) + policy.load_state_dict(torch.load(policy_file)) + policy.eval() elif method == 'gail': - policy = sb3.PPO.load(policy_file) - raise NotImplementedError + policy = SetPolicy(env.action_space.shape[-1]) + policy(torch.zeros(env.observation_space.shape)) + policy = ReparamPolicy(policy) + policy.load_state_dict(torch.load(policy_file)) + policy.eval() + elif method == 'gail-ppo': + policy = SetPolicy(env.action_space.shape[-1]) + policy.load_state_dict(torch.load(policy_file)) + policy.eval() elif method == 'rail': raise NotImplementedError + elif method == 'ogail': + policy = SetDiscretePolicy(env.action_space.n) + policy(torch.zeros(env.observation_space.shape)) + policy = ReparamPolicy(policy) + policy.load_state_dict(torch.load(policy_file)) + policy.eval() + elif method == 'ogail-ppo': + policy = SetDiscretePolicy(env.action_space.n) + policy.load_state_dict(torch.load(policy_file)) + policy.eval() elif method == 'sgail': policy = sb3.PPO.load(policy_file) raise NotImplementedError @@ -158,6 +180,7 @@ def evaluate_policy(locations:List[Tuple[int,int]], """ envs_dict = dict(intersim.envs.intersimple.__dict__) envs_dict.update(dict(options_envs.__dict__)) + envs_dict.update(dict(options_envs2.__dict__)) policy_metrics = [None]* len(locations) # iterate through vehicles diff --git a/src/evaluation/evaluation.py b/src/evaluation/evaluation.py index 3374c71..3682759 100644 --- a/src/evaluation/evaluation.py +++ b/src/evaluation/evaluation.py @@ -6,6 +6,7 @@ from typing import Callable, Dict, Optional import os import pickle from tqdm import tqdm +from src.gail2.envs import OptionsEnv class IntersimpleEvaluation: """ @@ -34,6 +35,7 @@ class IntersimpleEvaluation: self.env = eval_env self.n_episodes = eval_env.nv self.use_pbar = use_pbar + self.is_options_env = isinstance(self.env, OptionsEnv) # metrics present on every step of every episode self.metric_keys_all = ['v_all', 'a_all', 'col_all'] @@ -85,11 +87,14 @@ class IntersimpleEvaluation: if self.use_pbar: self.pbar = tqdm(total=self.n_episodes) + if self.is_options_env: + print('Evaluating an options environment') + evaluate_policy( policy, self.env, n_eval_episodes=self.n_episodes, - callback=self.evaluate_policy_callback, + callback=self.evaluate_options_policy_callback if self.is_options_env else self.evaluate_policy_callback, return_episode_rewards=False ) if self.use_pbar: @@ -100,6 +105,13 @@ class IntersimpleEvaluation: self.save(filestr) return self._metrics + def evaluate_options_policy_callback(self, local_vars, global_vars): + infos = local_vars['info']['ll']['infos'] + dones = local_vars['info']['ll']['env_done'] + agents = [info['agent'] for info in infos] + for info, done, agent in zip(infos, dones, agents): + self.eval_policy_step(info, done, agent) + def evaluate_policy_callback(self, local_vars, global_vars): """ Callback run in evaluate_policy after taking an action and receiving an observation @@ -112,6 +124,9 @@ class IntersimpleEvaluation: env = local_vars['env'].envs[venv_i] assert isinstance(env, Intersimple) + self.eval_policy_step(info, done, _agent) + + def eval_policy_step(self, info, done, _agent): # Increase collision counter if episode terminated with a collision self._metrics['v_all'][_agent].append(info['prev_state'][_agent,2].item()) self._metrics['a_all'][_agent].append(info['action_taken'][_agent,0].item()) diff --git a/src/gail2/envs.py b/src/gail2/envs.py index 9a48f01..b9cef2b 100644 --- a/src/gail2/envs.py +++ b/src/gail2/envs.py @@ -1,6 +1,6 @@ import gym import numpy as np -from wrappers import Setobs, TransformObservation +from src.gail2.wrappers import Wrapper, Setobs, TransformObservation from intersim.envs import IntersimpleLidarFlatIncrementingAgent obs_min = np.array([ @@ -30,7 +30,7 @@ def NormalizedOptionsEvalEnv(**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): +class OptionsEnv(Wrapper): def __init__(self, env, options): super().__init__(env) From c5e68ca33afb649ab8d84fe35bc6be3d780f3b31 Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Tue, 15 Feb 2022 22:01:28 +0100 Subject: [PATCH 03/10] Add model checkpoint --- checkpoints/gail-options-setobs2-15-02-2022.pt | Bin 0 -> 15019 bytes evaluate_models.sh | 2 +- 2 files changed, 1 insertion(+), 1 deletion(-) create mode 100644 checkpoints/gail-options-setobs2-15-02-2022.pt diff --git a/checkpoints/gail-options-setobs2-15-02-2022.pt b/checkpoints/gail-options-setobs2-15-02-2022.pt new file mode 100644 index 0000000000000000000000000000000000000000..76f70affaece3b8c27dd16f790999eaea4229729 GIT binary patch literal 15019 zcma*O2{=~Y*FTKR5>Ye>X`(@9an`z#C__m~C?TYhlA(#B6iE^?M1vtI86pjwz3!+a ziUw(-G-(i((y03TKELOE-sktc&;PpKeO>!L_t|T$ea=4n+UKnOS?l9uKS*3eL|R(p z|5Fr2ltjEYtO!^Y?x*YP9qK)O{hB}JR;*YT80fbmbk(|`5ZU<~eEl}~ z`C6}95h@!jqU{_kxLzYMmh}N$24GZ)O9u^rqJm5bX z|LNvmjnjjIf>k1eRsYp4_ut#6&j=n789efTljsGHiVPn8ze)6i$3zCJ{cjS3;IWax z>i?U>FjylpSo1$5VIe-#o$RHg-+f6oNFR7c|5T^{>ogA$nf7nftUJTTO8b9I^1sda zzn`f14YqhO@+`^sq|r9VNhmYrGemd3padIJw0zu%0f+xE_mOT)V`LH^uMvrperKVN zv@HMdnVlcuj=Aw zZF5Gs1s7OEgaZ6b9gin&Mqs;o8a}StM4y^k*};H6FnGjjfm;1_KEP-p>A!4*LysfT zuB(7-wJYH8ZB1ZT=F*B~0(Lb20DK-bj=J8@6S&V1!}+deX!vC?#@L_2Gh90T=JHYh zyNKXgL?}MZsH59^6sf%ME=Y(91Vz_pVT{BqHc9z4zvHJe|K?{AO}JVL#cpHS3m-2w z+fRWGduh`B?Jt@4=siNwDbY-4>{AvNVoQpFUECth9E;Arr@PlAguX@o6uhPl*3^ff zRe2=ArBb2ABLmQOj)%bp#_U(f1L8;Y;SPuEBsu32ooaTY%jE~rYGMsvVjW52ro7;0 zN!{Td59eWJvOB1~kVfVuMv-l6are#VaA3kFqJQAfx zngS@O^QZhZS=hT+8f0$+#XW6gipL6By6YXXxSc~ylFgu1npNwlR#gu-Gjg37N` zG~uqQaMr$b3~W-T$nU;%z+a6PDsO-(J4OlOWy*2i>{amU-b9++GfFrpssny~H53?6 zF{K9%y3lJQiCc3l3AFw@QDNVA{!4JMt*rrtl_BWT(&mW@ZFs;|XJFg4KYr{D#d?-3pOZV7RC#ax8A6G1di{tB-Ba0oj=6P?Rvll@j%Omxvg zKmTZmY8wrUFN+G2@5G>4YcwhTxP*R+Lort-fnw)8WhW-?M3<5FbZPdVNQMu=JY0#{+8t&TfP{|z3)-p(PZot z^~a#FD2fV>MSsm}+*}iFY&e<6NqyN(e-Euj=fCIpALkE4Cfh5N4VuP|^{#|!-&1f* z*%CGH8Dn9&7uZ?W(6E>9xfAORp|avL>O49_Rc#Rhy=;4I%2)$eNdY6K)w1lb-=OFE z8FH8?PcdpA%-xJzIKx;q!SmM-$#>~q;ng37q^$WJ)6yQZBa5NGB8`hs$0^3|@Yio`3|(e|vfmcL zr7g0;?oU%_cZWV&u3t+Z|FkjXHImreoXL3#Rp?hpDqeXM&Grd~Lte>FI(K|DTNmsB zI)O_Bee=$<-NQ^V_JK9aocyXPBW?$l%jG~F#OGy6gDlphl4uIgu|6* zq4SVK%+vlj?9fbrxivai^nxc-rAI`%Yw7LmA!uB*kGaJ^VsrKl6^^_*kDdjo31<55 zq&mkYF3Q;v)Zr&LK4lmkNa<&b+pE!H>JY(Qs}JnGe!lRqBddAM+2&;_3JG zLe#n2PZuW~rn!^Wf>Xp92Abl6GUpsBR7n-egcZOvJ#TulvzhZRPh`bKA6T##F+KG@ z3R`JNw-&F&tMOkkD?6S79JgbPqczyB_b0vUi^=>)F!p$h6Bo^&K#hOLJ7f;pb#NbnGVPrhj>I`9$rFyjJ zb*#YC_bu95=VHFs4B9BSL-=`RB?KCM5qb&M;L(&e_HE?>8a2X$JOmZEsnCe#9Mcxg zC@o{P$*ySeW;2YMb(4iY$q-C-98YamImn&cj$Z>0z=)}-+-?tZT=u{XRXvX3UP~|X z%#js3wb@qpC`JgCN}n@X7cHhMx|=l@ZDc<+ZlSi?S}N=OL}?}pbo})x{3^(0xf43r z+59*`_nW2UIA%4qH3gz$#vzQbc+ULZIN)frD7Mh>9aKpu)5~Hl+_5JQgSUxe$*p20 zr@Wlbr>;Ov|4gd!0^Bxc0Y#iHrKyWXGCkpV%Hiecd3Z2e;=r((@1%Dbsn8HM5?@*F zA=&F5xb9tyFubpq`bwiA0E&FW=R&DEYxo6t$2f^K{4W5=fI&~?Sxc;<-~J5UzI96fK+_*b(r(WV(%7taE% zsUNF*cE!>q)e`(J?MeG$E`y42FM1l)Q-kpW>_2dv7Imbvna`1)7Hu*=zU>t2J@AxS zdkhx@7D+K_X3cJ`*2B~fYS?{i0a?EC#@X9D;AqS~vdz2Aoj2b`Plmn1H`&kVbmMw@ zsV~In*J3P8+D15bkuN;kP=S~F1cL5jS#;1+qpK-DnBeR+a(OkHTxRz;H}6HQxy7{J7pdp*AnJD2VMlU`K%zbkP0W(0YD+(-^JkWD*Z$2| z6@K4*#rG&QdKHe{S57kbhtEJpV+S@q%BG+wC1J1dCCyfg6(}msV~$rw2yV>VNq>!= zK*-xm=+r$!plRbqf%m%D;^z<91dGe`aq9Toebfvt3!xs(z4#oZF9Yr2=V;Ulr8r8nelA z=5)Asu%MP#!KyW(uvj(w5BH9;2{Lb)&*om1x3Gg9Ivv36`LY!* zOrDHoev@GMlzOmyWy#rT`GK9{EvBUI03Ld4;q=}o=7}N4*(x^$eD~EDmz1UQkIo**; ztIMomwdw`8+gON3+YO8}RMChY(AS%tu~> z&DoO3pVIK5_s_QBfeWVW{@4;cq`ManiXGt$t1q&l{a^TNpWN75(LkDOnMK(NlkmRN zaO&EwhOE1R-F*lwZt8w4IgyCliyyMDzSGQ7PWgb5q#rw_!XRe#RyMwNGc4WvnV%zn zh|Z*);`_!LqG?G0zVz$@ON)c7woja`JJHXQ97oY~HG7DZwj(9mli(h2%BC700#Ug= z)n(sA*)H1&B+n*L&P7R@X79uJZS$}W#$xM6Wh(2oWM7LdXkAtq7^JFm&n)Bc!MO+= z)u~Bb*Lm(|{8>2Fdi?tC33I$Ejh07K*rzd(@YMS(1}@aa-lc{t z{7M8of?UqwO&Hz$QAu<1N04}2N_AntY_K(-1H&d>VV~z7#d6+{Pora5u&QN9|~%V(0y#&&qpFbeyXN72Kd=!gIlWclsk6!D9l^Q zvvCiq`AvFf@yYqyTtsdIC(i4k-K^cr?Pd^G?+wL4k@xsjml9}u#dw;~G8&$4oI?uj zO>lW$KK3rNrqr9++`&j$I@Dgy_lyvwwZ$fowQ&e7`LY0<&Ni|4$sd{c#j`9~x(h;r z3;FkVGuTMEm&~Q{2Xh|dhKlkd@wlfVb4lpr9p7JN%bkC-y6;8sS4SW7ofS~wxicGm z*$j@PPJyp$SM%>r%FV|o+&tDqYIDZLtv@ezg zeQt!#I4d^!+boc<6~~-w{oJ;eFIACa^>B0PXu31_2qYG~fOey=O!v@N^G78E>!GBC z{Tc8jvl2z|hM5Xq*x|skHZ;Q8@WWsj_Y3z3)y<lHFM;o>@6U7UdS54W&&7w52bIVr$0r?CFk-Bh^k9Hu;Vro%d&@YPzD zEnhL0D(zay_p%~?nDWFDw+`ZpOOs%ZWD>;L&IgyqDfmXEpXp!kLf-}L1P6Q-Z2xM&-=hcS$$gC(8uaZ-onCPQZRVvVz`oT zg-?&ivF&RPvp<$fm=-)AmTX=?_xL{6+H!!M8(Kt@qBr4zBfY$8Pb2sU9%y^ftgVjdWA2YIX;1-Ja-S4I`?2z;+(Nuc(xs$KEDHASI4ry-+^l9 z90ZLfO?+v62pR(unM-*y{d9_fKK#mU(cc8&mm6Vv{A3pUZWn%;wi&v{O`-8uQ<=Zh zHK@vqrA-oZ*_F$BWYSy3=R2H(;9?Vsad2W6^`67M!QNQXP!D@o zW{frceKV76x0OIeUo}iBbVmP-b$lAm8-9su2G)*Bn}i_f-mB_g0|2Hi)dPX=kgh zedSdb7LsAUCnZND!>zhd{NyakxO)QT*yTal>J#XYW--j{-T~e@lHj#DiAi^fqQT%G zO3lAWl?}HcTD?1H0d-lV&(mEDpY)|DR<_S!vb8dF zP-z6+R@zIw@8jS?^-xTkwE;}aS8*nBV))2w4lVy~NGc;1qnnl$Y>F%a|3QVoj?04I zgMH*3bRU*HJ;xr6-^~O$i{P>USG>5_2-FuUvdfEHFj8~}76j;Ht8qH@&2Og-;XI7I za|ymG+2O9dr`5(~cD&qs-|FP;)2LY1gjwgEVJ&yXz)raqy5|m}nD?9Tc8;0P)Y@xZmG_QH%ys!SnqJl@vN#zxr~ zX0=oc+2iS~Y!%}Yo~U7jjS|_s@8-6`8FnIj2N>H~u~nf{ac->w%Ua~mkC`f9hmOrf z!8~bFonuXYW6qgR{KF9)-Yhb01CZ*TZvXz@sF?U5Qn_R7n7M;me zQ4-7Ppt&8*tOV4silQ@u(IlI@l^?l&KWZ8OWey8n$xXV9U3dYZt5k4h z$z^u*`7ZdXEJX*tm(qw8^(gt_G3%K83%KVpxHGk!S$xr@hVK#JD3S`LIxqRz6}osl zUl&Ij*7DZz=Robw6p*Rf4}}fx%+Fu~B}XaKmhKna*w=ONOjj1qzB#~WPu4L1wPqxy zzZuQnO;aH8%08|;L6kOlm4bHkQ8?$8&Ye4`#?-kvY|q&w=5DRT9Q4!qoY7h+=YN!` zqB%wn-_A1p+t}&#In~Wl#8${FquFFXmazILud?eeB&W^9?68#}&KA+KEBBxY`*^3) z^AO{?o>hA8C9`%@`nD{PnT&k`3;PGr;OWD$M0Fj08}kcBSyX_eM+b9nPU0S73{;v+ z(IlsQ*1hZkpP!M%vj06 z9Nn>G0)TQ)EOW7dCi;Ypy4(x^fkf ze@})gS=g!7%gx?c!e$h-NdS+Qtvb{gufmcqNv2TVJ3E-7(+{V`&zM)_SkZ*~ZFuvDC}@9-r$=u;f$&)l z7&`6-JMnxLPYryiXb~esDtH03S}(C6#5vDKGgXjWo(9UkfFS z-E4tQGZb)R(MoWd9f|Y3%&9Rel!7Pv;M2xcta0}OOy3(0@3ZcjE31qmi{s1Sc~Ufc z-t!ySH9b5~U`i?(TPW|27}{>?WOXmb(d*0stXtC;&-mBit7!_X^JXFUS<{tfi(FxC ztt05fA!!ml=}cMKNo-4bEc}g<5gi?lBd2onHnwIyT~| z-RD^5%(2yfw>6Q|i`gi3RSMg!hoRfRzC>2@0)NQLAFgZvhWEoAX{A90W70=xa>8ZI zNMB2ukcSr|Sxn|TcpY9p7MBkC3H+n+tM_IOSe-bD5>K^p2-pm zj85c-%{`%uLW|p2#rwe?+vtI?fhSoy;fuHcBUySs|`!Sg`dDPhxMX%a-Qo@d% z^de0ZOB|EPFTRsI)h;B2!3a8-q+_0bQ*0#Va6!QMLJ}{m!y~Wu)5Km;jO6Zs z_n!@TO{#!z@nl@g!%=Kf@CR0*ATPY}WeKdBF&$oPPGI})8qxepc`B|xh$D`+V(sa0 zmh<2YJl?1Vqa;_5+4TmNW>^4;hJF|-lTT97U-;kF9zjE^H{od0mvH7wJVv|apx^EgmXvjlG*hN9b43I8 z>U0mAf8Y@Mxkw8(KR887yN03Rn^LM2jpy92$1;P>;b=Q8frWLvGwa!QkNNbgkbCTY zs#aPwu;wqZt!MYKaVrM1M=?M;B_pAt$PP=sJp(@n4t?I7WUKdFV7H&oq_CIH==Go( zg8m#~foAsRlFIIQSlfqwHn+3igRG{h;kMX+n#Q+5^vm}p>xO5?{;=B4Sh zxL=1$_6S9YO~kC9JMqK&@4S+`6islu%+6SBXLY7W*_ZxHRVB}l;$GjGv|rR5Z_krq zDI*SWB~54HxrZw5_ZNa|Kr(desttUbQh_UfhqtfQf&Y{(tT1N|+A=G=TOEiS3!Q1o zxH_62uE$<}okUz^IJD|c;cC=7Aw)V6B#$}`V0;R=(_fn2m3B$d-KO zV0zHJ?}j-|EbX8*!&4|@;y`>s(~fQq8Ah&>*MODn!q83HxK*)^U8s9lJ?L5;PBLC+ zUTIN8*>sgwB~-%LEpK4jwj>;57K_sjze2{QDsUG+1~vnBPHWFY3wM~#U3?3Qvb(XtT-?X1d-e>ZmqpyuWesVee@H$D)Ez~i4gn`+a zS#SB1J)t=I>t1Z!xE#hT_-<~r;Sc1_Eaa9BQO8whjF`B1E;W=!G0$V-WSbR$TFF)1 zkKMBDV5k8^Rc|Cut^rJD$yM9kHm6T}VsU-?Sc+_%3q!({Fy-DXQm;&4lO`TQna3gI zTYZtj+#L984`T37bO^n1c80A8p8sfpSyF3xQ}2r?vZ$DLl@?R; z-K+4v@&t|il#WU(zC)qpVf^f4OfUOWvFv1sxp32aaCkKugQxZ4{=riCV_rTt@=7uq zol3^&%}RJu_c61-OSEh8OPJ@nfSyRmqSO965UQn&qq_3=3z0*qZ*VS`o}SKvV!Zf> z_d9Vy_8=e;6oeFn9vUe zIXZMM{&U z;<-2$3Xi{z#W_Xw%*lQnZ4%Rl2{VUN*aan4tDJ!Kv%{EKS1}m3yFl5;n{0o814X_( zfu=vK*fsaN)c#xpM$A*h49i*Mep-)p?pL8>?WXL@u@O{1R0i*?AAy}UTjAlb3Fu#2 zK%)BHWY-~!0Xr9vOQshJZ;S&R6N9b$+$djcCTb@AV*Q)mRe$vJGmlbbC5$vKvNo=RamW#-E{Dp%|V{QG`s-Ziv*HgdHFD!A$gF?H)tv zm+U^;R-=X0V{Nz@UtP_|ymja8vloKx zEJ^lyUK7;JRzSI5yTK!|4K~CwkY6GKYKOd`>C`3m!=f7We{N&GWula?m5=`_~vT`O)t4^S$&Zr(V|53zpaC^Hv5?RI5p~f zAp&;iKQjx(8WOxcgoZlf@XLLB{3P=TMtHlDvHAIdoSS`+c1c8d$>J1kbl$}>EFEx> z!vhv=af5>Dbm`*M1p|I^EJO%OQFTTbuF+h^^tKe>9u1g-NDIQljrS^I0m*Eb;h7cGf3~MVqu@;F>es*zG%E;>}9%F3y8b!Ex2QObzHx zNH}}A>=N+fU0B_oVXW}wY!vJeheyBe@LHB4^fdcDOiYranx}TOx*-nTo*2;czdW4p zR>a$rde{aNUnY^)17CjX^S)n3Qg3__=OO;cqdn0}8ex=)cu zSC>8@C!QS}*JAqM*8nzR6-Q_C&#<`?^>EzoGIv>71GvvJ<~Fzj{_Z=^a?4EE$GV|7 z&N`Y*99>}7rF7{?+hl{ z$GwoJ(O((Rtwd%in#In%Q&K(+ik}De3Xs+|^0mo8Rpr$#FcGqTcTK0=soWKaq zbr!Rsd)xS9IqB$b*<8H~q;N&TTX>SO3!5zcVA0!`XgBR9e7CsIM_Ndt%&$ybcCQ~4 zOAk_F6H$RyCR;fFCjU#WoZEA7Cx*$0V6NE&Ec;Pm{`##MF4)@)7K+AHR(KWyY8SAk zMb`X*8FKK^0a?@v5!!NjK8}31lbV~8XSKGvUxeae$4GFAb7j+Ic#NOhX~LlON% z{PD-3t-S8U5!Eq4nfTINk`fi((!yOFE|^=ytoXUGxi^%}6AI{Sk`YZ^l>&QjC9#}e zSS+eAY8>HrsDpffu%%;LjMQ(SyK)*Ha#B74k0-bev{gmyMU z`VyBnI*d9D?{J^mf=KLW8|8i7MKZE(j5d#>1f6a4+;=i2&+^1<=Xv~>!fANdUJP#y zdkK?0;|C_bhOMnv8599IHAL*G<-812v zZwIr9_{=hY_f(6^MT6JXSlm*(n&~^3;CnrBTpnBikI%YeOT82|49cYz@ltqev49?g zyoI?DS&&f9p+!R`n>^SX4D1~-;ebD-daY%O`-0FT-5=iXIe~``yyd5^sA6?eRg~#l z3|il{fkl!2U$me^CQM{t>ZW;oGdmdBRm?_jAir@(V@0=`#O$EllUpjUGc zxb3Z=zv%&ZCQpXEQ$2C%BstWyR>A1fqcqvt#k^pXAA7M%8}FwMcn9a}K+$s-l$*@} zxo<<6=cWu;`FJ{--ekmF=*=*v?qLS;FoS!>w zS2{xZN7m7<@dLRZzei%YyDOIT>Txz{8mLiU#OXC~RC0VgUWf~0Ju-3wnA8ankLLMv zKey9r^-8Y!;B1JoY61JjQV?@)326MZ1Fx+`tm3ybAO6XaFVsE@i+ZoXm5{wGAoMc? z*nVR7~!LG|iu#VjTw)l^)W=IjUbRqd= z6z3zen9gPXfuFrk*^Z0q7@(d8zY^Q|zj~8MBRdRk%@bkWu|J@0!6Ej&F2;PG*989Y zl-10dvUjsqk7)jNcp|w5 z_4A3%Hf&JEI7*DZ&;AAursX@2qxvv$Y#5~sIY$mL@R4U#(=2JutWbI`lgazt7s2!) z+U(wqk^G~A>0r;8X#+T=C@)PfYJ&81S_xm!p zdx@<3S0WsBk-=k^m($|T3b0<#!s3To(5u=Gx*5BN{SA}CiIMARL#vPiT^G`-%e(Mz z^f&ZPQ)0O>b*yT7I!*t&lm6_RN?|o&WHQnLcAk9AmNkwBspu#0-sCu0JTbVqIk$`D zzY(%0?+~0~1~964q$j==n4a7Px56bs!=#7YE<~g8!$?ka#9wUHGQ%mgx7qlF=~VRV z20OaNjr1e$(YgzneBn_wfz!`z#0gsQM$&1z8lj3|8)LEAZ@ir{xZ7J`(eJ%}E6Cq)*82i>O0*@Wi;cAu) zH>^z#n?!(`y)+m%*nvWw^{MgtWi~cWgS;fSu<92l;EIGXv_G3y?KNBljXxa(;c-9i z!axq+2Qz)}|MCW^i(a!U4?nRtX_~xbgD*ev^lK2fE`TG??pC*tKgCp%teI!@Uw(^R z0W)>k$t^3s%{7VT!rZTHEO1C1<>gOEW1c&6w0aPDR;0 zkN6|Y=5v9!JV4>AlHl;L_2?y90i0zm_N2vwxach0w%{h6Xv?D1gAv@i6hoMr+yZ+{ z5Qhy)G2i7ajlRbhkdN$lP@b+p-_1AEg~UYz*wPU$ZuL%dkgsAdmWkn}I|($9w}v7r z@dKFUJ*pUPfkW0?W4BETy<4Hl7w&JOslBQseX5aO5u^&2P9TbK-;zTzyD{AuE(x%4XJFy*>Rp`K|!9`FjpCovq# z+_J!LXFA!vF|KrY!CrPRX%kKM`AiAMZk%<(V)03dV^w8c9TQg-aD&slO z;`vO!rji}J<_moWr#bz}IWTXS8^8Zb1N$Z{gNl+gyziaERi?G@s~4`N;*Xr*jPF5V zaa|D}v*`fIPH`~*GDzTl=OW4#WC=S)&lHxsZYQH_7Ra&^1rI9yg`w?QLNPB-;d(a- z;pY;6;h--w&_dM#bsHsw-%(%qA*d2lXHKKJ4~Amg&N6%|UnrP9Wscy}?L^_}Km|ct z*le;}o+$Kj*dcT%)E5SFR+RrC57Le6g>U05g<%>l6l79K4dc5JWQPdBCrPNLI#sy+ zt{zT2TZf&;TQGFrNA5;eBBnG32tyaYV)oiY1x?c)(^u1x_^sm>w%{6+(!7dSa=Xkk zz14+3SJv};0-muEc_H}rMVcVwg;pr(}6e(Crf4b(;#A=N^4p;KT6nh)BICxiOY-)}#DJ|Zg!+bAi#Q+=9D zcQx{Y*PqEb=O-)PGut zO-??zyM77ExQ}C3Uc6#9Z#cn?n&Z_T7mTTJwI=^F?FCt_3mC`+93>ESIEGFwRj9Q) zAD8B9kxD<$M7{)}$>^mli~okZMXj-M_G)2c?oK?HSO7OqnhMr>)PQSV0os^!F&%#i zX#AxrSY|RBZGvP4{RWdT8s1>Rwlo@|t1C49+|S+*7ZvD=juVzdiU>6WZ(-?+@0dJp zkibV_D!~D0_{P!;YI8%I&n!&@J3=Oe&V92?nxr7A9;+fwk}EgEU-;g7SEQEXCyVA7mHf}P6h)MWFNQdSIQ zU&NLRmsh?dhw+LuWBodTX%(XpLO_G`59oYKMvz>jEKroH!FG+6xcQzAJUnv%*Y22r zQyixWa*cf%wWP2Qvbm@e6oGeDhq3%8!-eXTvxI}}{*c6smHceK^#YH!b9ig9v~b3h ztJuAuieCFH5ZZnT7W&;0VXd}#^kzsu4%^T0lYyv^R!R!DJTIa^sr~qCK@}$MUV&8} zewZ~%PHn`28(5j@Nq*MuJ%4JjN3d}sHE;9^bGhxst3-|t)~NAqxxFhd01D_ zu9GE@=+dV*)>){kJ0D(0MhMip)aj5*C7vw31Wf%K1uY1|={>6jzj|EZ(n}S=3)>L_ zZJi99+tEvMgSzST(fxSfXek=+dW>KEYB7v^IN;Ctp@f#0z_J`+)Z08NytHosXOogn z{Z(IIZ!9A6PyN3W<`hKC|BEp9pDaNeE0h2KjJbd3&l$-6bLHDdRxT17;0^v$o$N*b zIs6wV^S{&h6#o(YTlt?1&Hrir-^VoM-z;Ab&;tLvweF1nY5U*rO5xvZN6h?xZU6Hb zIN3`|PW?}%v4f`lm-;{C%zw)Nb3FgAQ<3C>F8(P4cK>iaMdq6Q^D83q52@41Ug96h P<$%o~k$?36rTc#XiF3P} literal 0 HcmV?d00001 diff --git a/evaluate_models.sh b/evaluate_models.sh index f1a4b62..830752f 100755 --- a/evaluate_models.sh +++ b/evaluate_models.sh @@ -13,4 +13,4 @@ python -m src.eval_main # idm python -m src.eval_main --method=idm -python -m src.eval_main --method=ogail --policy_file='checkpoints/gail-options-setobs2.pt' --env='NormalizedOptionsEvalEnv' +python -m src.eval_main --method=ogail --policy_file='checkpoints/gail-options-setobs2-15-02-2022.pt' --env='NormalizedOptionsEvalEnv' From b78f95bab54f2430c4a885c217be9e34114a0224 Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Wed, 16 Feb 2022 10:19:50 +0100 Subject: [PATCH 04/10] More checkpoints, adjustments for collision check --- checkpoints/bc-intersimple-setobs2.pt | Bin 0 -> 15383 bytes .../gail-options-setobs2-Feb15_18-49-05.pt | Bin 0 -> 15019 bytes .../gail-ppo-options-setobs2-Feb15_22-05-38.pt | Bin 0 -> 16087 bytes .../wgail-options-setobs2-Feb16_01-06-27.pt | Bin 0 -> 14763 bytes .../wgail-ppo-options-setobs2-Feb16_04-02-56.pt | Bin 0 -> 15895 bytes evaluate_models.sh | 9 ++++++++- src/evaluation/evaluation.py | 2 +- src/gail2/envs.py | 13 ++++++++++--- 8 files changed, 19 insertions(+), 5 deletions(-) create mode 100644 checkpoints/bc-intersimple-setobs2.pt create mode 100644 checkpoints/gail-options-setobs2-Feb15_18-49-05.pt create mode 100644 checkpoints/gail-ppo-options-setobs2-Feb15_22-05-38.pt create mode 100644 checkpoints/wgail-options-setobs2-Feb16_01-06-27.pt create mode 100644 checkpoints/wgail-ppo-options-setobs2-Feb16_04-02-56.pt diff --git a/checkpoints/bc-intersimple-setobs2.pt b/checkpoints/bc-intersimple-setobs2.pt new file mode 100644 index 0000000000000000000000000000000000000000..944cd705c479bc17f112edd03f109b7080d9d275 GIT binary patch literal 15383 zcmbum30RI@*FT&F4Vt7hPn3w}L3N#Lg$5;(M43sYWQxe#B+{rP2^rEP8AElQYljqN zicA@blsRLRss4HH_j$hC{T=W79q<3!$FZ;Ty7sw$zrD^qoY&fGujS<-CN3f(B_;Cz zN{S*%BJ-Bd4-Q!sXcjOpe4gpD&_%9eKFL;d>oQZHFzYHk`47_uNZJWOW`{Ll&)&nX^35e`Uy` zfH~oTOI9phK4+DsOqkfvp?;FRuob;E=Pd{f6aNb$l{>p%P4(4H!6Z}NO!=xs8h!^3pE!*u^Go~7x(#MAqO&?-#- z?}USXi#Pbc#54FUp5b3yL&C$1yu<#}{nq~$Z`dDf)?vnfXB+Mj5FR$-H=C$OuXHAV zNoV>Oqgi;Exp$bwzooPOOFGLx_-w+g{?2FZ(K{G6|Kf`av;8gI$iKMk!o%#n!yHy_ z2owB^Z_eVt@Oi(dSs3@Pfr6oay+b4th8w~h|1}j>1ukFFJH^6AdGuM~C;G493md&5 zY)tUK3VP@8zvft&(}pnTKZ-2ETsDNc{!wHZHg-dp+aE<%VeT8k#{E%b9p9_X>#Df@ zO)(VH{fvLApal8etI(`@Jo|O(4yQ9V7Il<0nR@j?CNTI6Vg2sV*tH2PuzU_i%-@9q z4-qDJa0MPU)x-Hp>*;=4BAjzQ2kMqHn54-x_Fd&9GZt?aOl?iVF7+gSLs>iBkwdBq z@WnwNez8#L2=qSx>z=kKJPh!Rt8U@a%B#4vhrmcUzo9&{2cyi zfE&}X4`;raO9awIV|4JkafJC10(xQVmIBJSZ(APHhS|z{^7FQV3)X;40E>O_3gRn z*cHjf_e%n~nYnl&Bp>n?Si$a<&&jWN7t6Y4glBi1Ws|bIm`v_sd=p{B%7ZE>cIF~n zjeyQdnJ67Np3RtI!KW(8g4l9LHYwB)n?9$2*JeL(RcX4ZQ8Vi26nWdJ-QgL)XdwhK~8lDPc z-jTali;g@~c)bo2H)dl;Nhp?U>tOD?Ci1?25Nkdc!nP&5p*^jW<&7AIZ#UVp?Me}x z${RBl+bP8;&z8@=(t=OD^>L~0YqqaPi3LTtsX+&oJrjJ3ema)P^|A@47?V_CE1R~Eh`q`A^s8<=kpDw9(hGu;+@#EPktcx zy_=lO$|*6*s51GMkb0iFP*j;2zjsw6dnqTx+3=OgUcSJNl%~S7j$(YYya>E{^iWLd z3^k9ICl78X+x7VYjr_8j?JR8GMZuXy3;y(N>nV9y-Z2Cq&?T&JG2B8}Y=M5C~?9xaFJ~mQ9qvP%J^~ z0XH%FdI(sg&&Mg>IrN&Yj@Q@4vn*$1KGtHDGp&ZR?7}phnehM&UM^)OSD!%cj)5$G zur6xmNRh|WHa@fS4b)u$rRM`l9ToVo-Y> z&u(!gtRdbBCU85^aF8MD4UJ=mcMhU`0uN~VIvSK_&*z&e!?FDEFM4Kg4<8y@aQ=sT zRQGWRhy?4P;uj-S-Ls}Q$Dw;rF|HaJimefr&@MZLTmNV$8`>uV zm6G(>>Btc{RwNU`YX_l(-xmB{9*qb6gf!h&4KwE}p_0-kR(>fCt3yLDN39V~uMKC< zC!J?KqZKMKGo9ZYQOqP|&TuAO@8F9KpmN$|wAhh?J5I^4pJO@dxZ#Uwi>9N;Mo;Yj z@-5h@_hBW5PH1#boH``*9c?b!Gk#7gt57e1)%!Bp+HoUr7OFTt+TqE5Y)oN$HgnA1 z-GTMnJ%>f!4Wkc5S7}xSp`_(OZ2H*`SGn(GJuRwurSS;celCGmnfewc`1VJyfYCIv zunctzl&H=;7jJZCW8{S;H~`hLGj7~T_lFz=Xm0q0Xj^kh?v{rjo4JM4BADs9GyxwVeK6S zR%9HH@|)C=7xQJGwx4BtlRvQZQ&Z{8f@IdbnX!ctDeTg^AuK%3kv;sd3MCB2U}42G zip*pv3jD^uG#BLEA=2gd}6b+e#Z8BXkJB%9&rBL>UDS9u>VPnr; zq665_K?r8`}9_J zMqwiwonFLjc3I=CR98xDPsO>f4??rR13R-7KxD{yn(zIMdZP9M`s`#^>jzO5ri0s| zLpUZ%3cM?3K<`pvxohg#eE;`gze2-tt=>2OxL*OvugiwUnZ~eZ&N7x?k;ct16ms$B zcfjIwM!ExZL6mCvMP{ zc(3>WG5NeX7M_#Doj)D%g3cwrNA(6seNRSR!4`I3G@Y8Gf?(ccy+ zl;J6_?Mz~D6h*(f&F4f-V;0gGI51S6O}l4x#9UdH5m-{}KSX0g1L<)jzu$euSFG2_|M_#>`A z{fM6m3+A;8O3rz*gSx}fT5<G+}L!-nU&iH@b*A=Y*klrXrKySpZ-dY3!T|>zw28?)W&Evoe;XvIeqfyA002 zXb_w_W57lYEMy5Ul30r2Q|SNNl%ge`(=Ih7cFuVY>FC2%I+kur@%V;ouIAtWvqF81EilBjFuLXlr!}JESXTi z;wCL4(PQh;PH!H59Xb^Y}6xhKv|ReKj4(ea$A?uy7n*ls-v2<;8K7wknRZc@EbLQz%1E6}5s3 zxV__w>7j-zW?e6$YhzNm@7v7jo1GXdO&fscSG92a28`qrztm9Pw|BJfs2>czb)6zh zPtdV*OQ?Bm1V2@FI$b{iv@~@R&1&l}(EN0tZtVX>ubO^9^MNC9%ufeD_DN)g6Fca< zNe;{k(#Nam6Vd6)H$JkVxkB||ER{M2>L-GEbGy|8=ZP6vISn4n75| zimt=?d-JJ1NsphXm23AP(ibjXzQw;!9fR$jQDAfS3{=eM!<2h2z|Y+?*d&1r?#!LW zPSh1bqecKu4)p?DGz997&1Emvn+Q}!mvIND%RttfS$OW*N-(@*ijxPV^X0;c7)XXF zET6=e%{dFJhFMbTV?E^EtEk_~vCt>UjDEF-(8^I_*mg{rZ4Q@X!&4@a*L+h-$kpn63CokX3RF{4M6gMUqwOoiF zF^u1TESG#S@0sVk^(%UV0}>oT{t-gybBLf(1!)waII44vEKy~ zXB-iDD4F0q!ExMgyAfJe$J6)^V(|6&9?+HS!_J&7!@g3YAuO5c$f-T^@uQs}A#v(<eHPgXs{Z@YSls;_U)MR**B+7Rf zi=sq;kW9RWu(mrjpy0L-Z};6tk27|XMyoxizMf+x35Fzp>=FoFF2k&djPm@KV))rj zoXwdSK!paOf)meoLF_&k?2{6P16<-Ep-CW^sN9#m8$on;4Uf}xR52k)9x9J$;p|xh z*)oF)3|tk+Ztgn@9&I9+8xe)e$Ce0M2M%YOV{*~K-c8^(L)y_r>LUJJJ&KHOHPWO= zPbS3S*kSt-YkX{3<>`lDJtP4aa{btNjhSrF$|+3IKL#5b)!EtI1KFHgSuA#D7n`9v zlzmEi4NK?OvE8=Ks37{0E{g`C{~%}nPFXX(*|-S}Kmm7YJ5t=15p0XjPn?$c1(f4g zf%)#`xUf|X0_Wc$iI6P#sN06dGz6(DhF1<+_#zAC~4$LTK`a%y^+~T z7hZJnhqMRxzPnV}5S0^jcZNR9`jkQ|3+(B5wLQ&0EJs{NDlT&y=`bZmo^4REhna1u z^d){bRL8xg&;!fiRGJOchCUN44;5wChZ&KC&Vp>UkFepSA-Skuql~Of zDk^-*Eq16NXszQ@-O4C-gd*E){)l`6&r#Tmk!<6pD9*OCm2w*8*wRt`V3b`u<;u=P zxdX=mK1EV-*9)@g`p)ItZs8v{TrV#=+?U<$mLi{%GEDqQ7mb!G1<%3pX|(X}vjmjS9wKME&DEQ9^8 zlyUw9AIyHB2j$^?F#cXKTq&&OCS_&|5`@E8keoQ4J)Q~1+lFAM)Pst?Mk@rjxBukC zx0yozmP2qGqd?N<12|juM}xtg6s~=i-18(L#p)n9s^-&(sa?Vp77F!H@< zrDw}sntO65smR@A?j5~*-48=5CvP3Zhdimpb?*<+ps8i7(R?odwcP?IXmqnt-%8mm zdDTiMQ>Dt86YS6;?Iz_9>yK^wx3W=L11q`w;dpLHAgWrUcsauHy|x5~-zoq_llypYrA+0ar(3D@Vje4Be;J$~Zibjw`jvZT zo#M}2-G@fK_C**psM7js1WWj~67$nXvE}Xxq+k+{BC8dkNf-;s$udk@eInebm_!O) z$()g&7&G5;LGb;~QTVA8!Eg2#VV74H@>2Wdn2MoU#ogT#S!11m{n&PZJ9uFe)qasc zn`y`S_?bH}$bA?aJ)8Kg_1)aQkA*bB)&j2Xn?uJA+;pfm&8L@z8jyIekRF|Hp!off z^xi25`Ov*WxgIUp=-~skSzCCeD_h~`s1%5)Gws#C&r$2BiMXfxG=1DQLlEO$3qPhV zf%#k0seXE1mz1=pI0uA!v@Z?(bPdrD}H}v7l9Rpg= zU55C_L3F3wjoCe~q2&5Wy#3h>*lu!%CY&6{eAXN$ee+3}YoSa=+a;iz#Xy+1BU)PN z)2bzjbXzuEIQ~T&i8jy0P`Umv-%F2mJ--gtW?rOw=^?B#=)<;dU?6wt4Eua}1)Nyj z0~jzI2zMQE{A1*ntr56o>^y=woPYM*{ zXkPbq4pi;5#y2DPK!>+P<)f9lSk%vsb`9sb_^EY#a?W*lbiaT<`SB|KNbtk~XZO zQrGi3znK}v2i%-gVuDPPwFI9tyxyl8-*=6mufD<<=ro=&8aeVFc#{y?!oGhoIs zLqXVtN3e5%1f&`!z(Cvge0f3~OUqjWADk?iuW=nnnWR!}19AF_9i-a32^O7RNnav9 zP|t)$sFku{dBumw?d2izGu#7hgJszKqcKzzJAo~4e?vxD0)CB`kkgu#MH)wTvIUC+ zVS%L!-bi^#6Bk)=#WgYf@^f|6pcc)R4!lD{B%IiCheY&^90(Ep>frP}#%-AR1WHfd zVo!p#Nq8!PN25~VrFXj^bn!IMv5w<@OMY_~-1LnCNNz7ZSRcNm7vtf$2NLUPD5 zL)X(r?Av5D)QIfE96u^EhrZis6nB_})(?d0dImt-kFWCe{n zr$M7-tkAeS7cOVH^RnHJoQ2|d*dx#&8`lQ9B>jsn4%iJI(LcF$p$cT=|BPJNHb9x( zlqx?KLhPS#NmFCk?yK9#UDFA+8L83T1UZ(nPK5I7fn{$|qjL&Z$!D}VS!NaT({p9f z@wq>VifA#Vb@G^*`VHFe%@u52q=07*1fZ_^D%$Z*0*B2y1)=G)!69xoSG4UCoStRP zYR=f=6df^G9eN86{ft4g0-iT*orW0`t02oO1#`E8!huOU`DaH_TCR_e95zD40ju%NKDy>H7Ypq<|Cef#k$jQS-}I6 z-`J3KzzN#dr-HM}bmhE_lDV?92P7LIgIet#lqgmrur2xt?c@*7#+jgYfHxk?7>ZsO zSJU0@dN^8O4O*TIpZrXP2f~-&{I- z>DMtI%y_*1wV4a4rCMJT{^utzXbfONHIOxR`eh)*>_Y@zkwx z9acq%QA5B9C>VmQ`r;Qz%wH|ccIBKvd9zK}yPp{m1cx1E zFfVqE@a(c*T)V@4!NmA3=GCM_SY(Qlk1bGf^$v25bAbu|h*?)gaul=|2dM{s7_v3W>*3m*q$K~0=(k6T~ z*PeCH(PDvLcYue)HQ0rhd95-jmNq1dn-cy83cvn>C3)*${Q0A?(vD7kuWc>EEV^(_V_T>3%PPe<`$ zg`y0*k$Fr##*d6E5X>r6MNOq5_`2jQDV`e-S&L2RGPj334NjBHqY`ny^#Eb#}Oc~nP}R68IlS%ej~)xq#@zSMl=CTY%U2j3qqFuzv0i4ZT~)sGB;y4R!ZZ~>8BxmdMkZ!{UBJ7sYjjB58$@(GOoSNk$W4v3InBOpuAX} zv~~2L&oLbuKe3SD6i+uj`ZJB0vJ7hGf$YpJaC=`AIml#z!L1muZ+uH*WRr($Cy@*d9Hftrb1M?~$?R5?c?m$XPGoaAmJO@YfUGB6BynUK4`^3pJMfx&|i~ zq>_t@6e*dO!QwG)EMGz!VoOEX{VI7}ygd*MS0-@{Q+oH-kypun;t}3-qa(X(-wa=7 zIigg;AxfI$jM{dWX@_f_pg^yYGLIYKYMb?(ZAmnAbe;m$Y5RDod$G_oBn4!OrjXY_ zB~ASL+ZP0_bJ`*HTp87` z$Y*!7rV2%04#IC|&O?Br5>v4Vr}|4j;6U6fc1EGTneFE5=YpWH`Kk$f2Mq%Gm@U(Xj@x)}nCWzwNG zZz$%sSJ4=^WLR6E%kCbyOX6p@(z}fFG&SmzpuXD$hc`>o@wgCNA|60sHG(uo=fcRB z_u!CS?-{bvD%vny7pLzZEf|?H0TX{r;YI6*gJPf|>IM#CM;m+v!viNnPTDCb%vGb2 z(n@$vTbflY^TO+mTI`cur?9}wjbAb0kkCN?6YO)7hnmk1VVS2HPLwW$@RvRG(Ps=Z zev(V}lOBVV%TTynJq2d8y(6DwEj)2JiLXwprJ}JK%x_B#bX;lVCcSHf<25^=*83H_ zDh^_)Rt_+3{VV>(vP5=z%QUz;rH&#E^|0jPGq6^9EQIT8Ftsy}w5r;`%S?xDh}TDH zlTP}nc3h~@l*-vYTML45(NMH+J$J*FLDgXV&kuYckpNafQcfVn=yy2A2f`W z=IWqp>3CAmdQ9aO(ZH&*Iloil__b?4%zA$YcHfDGiO-DrEf->7^8C;A^Hw2v#yQ#> zzVSq;;sgO7hvVBzE^wj45zZfcN79XVX}0Yq7$LKRvXu^T37^DRuE;Asum1xYYBUT2 z68zw^#uHll!~*AWTO6XEq`;y$6QNsN$UMzF`HGM&kTkEGySnojFS$rSN;77%dKr8E zhD#Kcwq?MXPjW1{p@drIcEc`7FQ)Az#U3572aDB@Y4pv(Xyn-lQvy!%`iaUcZnPpZ zNP58yt&(HwLrIu5$cT08862Y7iw`qX- zThw&cu4sPth&#}r$ZT_!VA`4e*f}|rZ9k(yQ$E-6Em6hrqRNZj8970InGEJ16yW_E z!`YfEkz7Wft<<0x1m0h-ki6XjG#UI79tg~t?#I3KJwy(Lyb_DDje$$wAA{}awNQJj z59W;-%%ta7vB?YiL+crBdQt33sxy3XK)4VJd(W^&?W{s?w{Xmx?FBtIk3i6cX3`~X zS}c7L!{6@*+?heLqM5?cCqzK@lo}?jn9Cem4Y5@39-N+yc)&c){^*=VG*dkqO|AM= zbV3~WPPqj1mDDMC)=|z^Ig3lk=)kY=xsmxeY`pNo^Z zz8G{))W|rxr&Gg6*Jxs+njJi>oDPQR9MoJpPR@8i;Jj!A+cruM z&zRNF#0mi?y3d(89^(m*$F6YFm1)%!nghH}m0G{b5K$b9SMu3wh` z)7$&8`hWxOUUr3U?p+DaYjyc{x05jBsy5u~T2D!56){s@23Bm+bC~Vm1r3Wn(z&+- znA!M&xW@f1b-sTN%bW+o>68vQZmZ1DXCRHf=foC-?c*)qIIyvA4wL1+PjGn_u(}u3 zgY)_AX|nTEiXn8N(H9n93@mC0O;@>CF1m zJ1%d}T{>T#%BsKarw$XYK%P%a0P;bx7cWS$X7Y z;KY0TeSn{J#>}R?9=sCGS^TZT^!3FQ@G2g~jBA#0cX2yLkCtYu%NKG>_Dm{c?uFhb-r><5u*~!!HUBm)B|Vb711+CRqNylttJ9FDUJ8z~XTS}xq;lJM9NwBu2ahnw zAApYeB~x)!jksg>TRS|sHkY+a=hK>{(_pXfGyIez$rP27F?Q7z;YU9s?02^qPCO5& zw7hkkTGsBs+pmT*$JNrPwBr+n1a4znPAmoKYh4s#cAwM^dV$^fGF~ai1(nan@fYZ5Le3v?0r~G8jIw3$%=rQMLbIthN;} zvxac`KJhiAyStIv%lBlf+C!~kCt!(qm2m3ZT-xz$C$BkU7b(5!5G;&1LT>klV}@=9 z92mL)e7)QG+vDs|=I|AeXkHD6CQ0BSl|k%)`53TTt&j1spScgSA42WnCR$J|4t~cC zSi{hpFlmV>Sm`=*P6uzoNTC>;E_#%--_?@3`3;Cs-$2vDH**5jQlJDL6cYldyq^oE zvY~iNOcWOmlx90xIF?&f4u>Av(X1*77UV3;c6-mJ(x?VO_pKSQwRjbXc6=bSb)#Wj z_gu=f$m2X#O~#O$5>y;Jmz`^JhNRu+1Y17EfN@DGpRhZd9AjqC0++Aw*)fI7jM-1) zzw2O8WhTGwkS*YXEU4Sv4DM}TcukrLFAk)z3tqi`l?(Th;+j_4dhH98JdUBhQ+-gn z*Eea=wS2mhqezRM7(>p{?R06TGjsHuz`0sf!L;e!Tvvh-y;7Y9%4KE{7nBdX#lKVc zm-7`uohmTCp@Cy;rI?;pIMoe`r=>k#xtg;l;I+aOSo%bf8Bbe8F@v|jX^jHW=6hO?oq!RCNwOnhxF2{S+r??I(_qeuTHv~tt%C!5Bo=R&aZW_!OxK6 zMtIRtUsvX~`52TLggeZ=c!n0vf5M&Ae@LU&25=*uj%6jgib!)_DF0kW2%Uk7EMeU+ zZr*VbSQTnV!o;<>dDsP7*%?i;`*nKvqB?XfFNHKVNaDRK6X3Q~6J;E);;*I)NZ( zH!Vrk!itgr+N6C=`0;25ncp0Vecns}$Gn@Qrg;OJB45H+{YS9WB@((W4B(BAe&F2; z6(YHB6|! zPd9C==#WVXq!k77FMque-Yl+$@CnFCH7cUPGIy*We3DlQ$Oezn*_@gn5_{YVNV0np zjQwnc`i~09Wc?nf_z*=plLn!8-ypU}RvhweW<&DoZ}`r2Jd3va!Ye2!vFqy(!KX$gG04W%JEF6rN@bno3qml$kTC+-;p~gb z9?7_*ITgB^9GOMx0(Rt*7>f0tQE56{h4fAvCVnuD*-2(|otj6uz2XNcLt_JH%OB#F zov4A#1P|PLtUm-l{tD;!77IMmZ&3dczsMX;&?&#~P`0^(?;6wCYoE;`8Q;Dvys`jH z&SXPOkrwznJmCscHbVLrA5=eH#r1qyhN&l3L(51P7;-5crY24SEA^oe)~(9Ag){k~ z5&3ZL!YP;>{8;$K>oIkuv`}}wHr8Cb2DUq1apMX%)5;GQX!Q1j1Tn{8PTvGNy4w}% zhPuJ%d1^Sc?J32VDl+edwVYJdHppIig4=s^7o9tu#!cI*k6j@_AnK>f?ceP|(##s` zwW>gJZv$U>>ox`LI0V=8kMgDCTVYC`7HbSFaaf!th1FlgnW64r$60E+thjgYzDC3u z$Io4kU{M4wZ=^znNRIAlzCV_6{v1@!WH2P?8?2f_%xM4^{OP; zwQE0oxEu^Yy?T&;j({DRu?)FQ4|x8u7Hr7*1tKHk*+$euYyW4^n%c_vy%sR~8XCLK(R6!3Zn<2m)=V{BS;5xyE?cR=uG8P zgW>s;2fS~%IyTGMu-v@i*l$5CKaM{N$_1O?@d-_q*P#lE2Za=`mkEn32jTS?DduLD z&RJSq;>*8?A-`-a-uxj!&A07QQ$!S79*a=%w5lDIZuUn_TV*^y%0+M}-4qkH z#<3+`UDWQD3oW)Exh~aLbXsIKQ`gAk9>h(;fI0Gbv=SV2LK8o6sA5r}_51!4? z;O?)k@73{U2-EvJbGv6%bF+HJ&|2R@*q1$#4%fu-Pb-j}PqzU(l{U`);2;(}`3lrd zlOfZKW~ia1ie?idXyYYW4EnI5{L%c65PZlRJ3>Y@*>w z*Ul+n5zYC}R$x}TUT|Am4%an31&@fyEF|YH_fGW_r6)T$_bs%L zZ#v9~8^c>q8p*ldZGhn#H{i}-Mr~%=IObLb;_2I&V~WrOuIDpdKi%JA}Vn@L5^C}!PRUEHuHPHNO>vS zqBssmq+AwAKeVEeLt=!A&)(5|ngO;}eOYNjCxou`feyti$`R%1!M#Ud)vuClj>Lex z`949;dP}@KsSp$*kHea(zMy?klg{7pp`{+uf>n2A*vFl~3h%6jk(>#}RNR5!eeE=U z7RT!HKT>snIUI>1xcqey@RUxs95vYk^?{sw@80W!t{C1+yGtM5s$zq;0Ap9m;#EN` zeOob%GCiHy*4L-Fxt}E&w={)C2^y%D*J7Uei}_CjzJMJ{;vtr&TG2>E)xip z9nG*o&l{8!XM^R+K}>gUk1+a&NrmI%Ztj-THITfki4aLq3RtUum?-~IOfUcZ?C3;#vy z|3ChV_(dEOdtS34(tAiZ+nR}ozk-F~$ALE*4qbQD@iD)Sd)Kv}6en$DAttdH>NOnR z?=fK~QnRUh?L6k`vK%FQ<6CSF=(E{5%BU2VjX6_?!nf>Xwnp>`?OL7-qlbxAvSw%a z6|e7ju5cHmEli`@hAxm!)Wa#MDOk`V0m;ujnfSv}&L@5~J68UHs(53ZaQP|a)^oH+ zZV#BBU4f^jdxLRg6#wS!ME2EWIt5>T06tDyIQCTyc+b&fc0-4edtf0ZU(a9zFHFYJ zg}SW6z!kxBKf89i4XY;Ep#QbQ{E1~-S=#ZlY^2p4>UEk!m*^gTaY7RAf0~I?{H<}8 zK#ckwc}MPjH?cc5yV$a!cFa7^0QZKv)2N&{EUxXQuE?!eu_1#QkADDq)}zpIM^LXn zOE!}!G=U($Z2TCs1BA81FuHsf_No2`zE2{sY}zm!@-Bd+wk=@dQUdm&(~%`chhcNT zZ7yxtR-9{WgTLp&^2-^)+j>I-f9Lh=n{EmtHoxMEz{J#7rrt|MXYl^>b{E_|{ z;rUPW-+ia$|3KP{iTvj|eh2QEng0{}cYoqPu^(jrfwlN2_V4e&Ke1}PKmRVZndLvR ze}51BiQRAW53JQcv45XCe`4?3{sU|MPwd~v|4-}_yMJJ9{)zqj2rB&1&wlp*5B9(2 zij literal 0 HcmV?d00001 diff --git a/checkpoints/gail-options-setobs2-Feb15_18-49-05.pt b/checkpoints/gail-options-setobs2-Feb15_18-49-05.pt new file mode 100644 index 0000000000000000000000000000000000000000..76f70affaece3b8c27dd16f790999eaea4229729 GIT binary patch literal 15019 zcma*O2{=~Y*FTKR5>Ye>X`(@9an`z#C__m~C?TYhlA(#B6iE^?M1vtI86pjwz3!+a ziUw(-G-(i((y03TKELOE-sktc&;PpKeO>!L_t|T$ea=4n+UKnOS?l9uKS*3eL|R(p z|5Fr2ltjEYtO!^Y?x*YP9qK)O{hB}JR;*YT80fbmbk(|`5ZU<~eEl}~ z`C6}95h@!jqU{_kxLzYMmh}N$24GZ)O9u^rqJm5bX z|LNvmjnjjIf>k1eRsYp4_ut#6&j=n789efTljsGHiVPn8ze)6i$3zCJ{cjS3;IWax z>i?U>FjylpSo1$5VIe-#o$RHg-+f6oNFR7c|5T^{>ogA$nf7nftUJTTO8b9I^1sda zzn`f14YqhO@+`^sq|r9VNhmYrGemd3padIJw0zu%0f+xE_mOT)V`LH^uMvrperKVN zv@HMdnVlcuj=Aw zZF5Gs1s7OEgaZ6b9gin&Mqs;o8a}StM4y^k*};H6FnGjjfm;1_KEP-p>A!4*LysfT zuB(7-wJYH8ZB1ZT=F*B~0(Lb20DK-bj=J8@6S&V1!}+deX!vC?#@L_2Gh90T=JHYh zyNKXgL?}MZsH59^6sf%ME=Y(91Vz_pVT{BqHc9z4zvHJe|K?{AO}JVL#cpHS3m-2w z+fRWGduh`B?Jt@4=siNwDbY-4>{AvNVoQpFUECth9E;Arr@PlAguX@o6uhPl*3^ff zRe2=ArBb2ABLmQOj)%bp#_U(f1L8;Y;SPuEBsu32ooaTY%jE~rYGMsvVjW52ro7;0 zN!{Td59eWJvOB1~kVfVuMv-l6are#VaA3kFqJQAfx zngS@O^QZhZS=hT+8f0$+#XW6gipL6By6YXXxSc~ylFgu1npNwlR#gu-Gjg37N` zG~uqQaMr$b3~W-T$nU;%z+a6PDsO-(J4OlOWy*2i>{amU-b9++GfFrpssny~H53?6 zF{K9%y3lJQiCc3l3AFw@QDNVA{!4JMt*rrtl_BWT(&mW@ZFs;|XJFg4KYr{D#d?-3pOZV7RC#ax8A6G1di{tB-Ba0oj=6P?Rvll@j%Omxvg zKmTZmY8wrUFN+G2@5G>4YcwhTxP*R+Lort-fnw)8WhW-?M3<5FbZPdVNQMu=JY0#{+8t&TfP{|z3)-p(PZot z^~a#FD2fV>MSsm}+*}iFY&e<6NqyN(e-Euj=fCIpALkE4Cfh5N4VuP|^{#|!-&1f* z*%CGH8Dn9&7uZ?W(6E>9xfAORp|avL>O49_Rc#Rhy=;4I%2)$eNdY6K)w1lb-=OFE z8FH8?PcdpA%-xJzIKx;q!SmM-$#>~q;ng37q^$WJ)6yQZBa5NGB8`hs$0^3|@Yio`3|(e|vfmcL zr7g0;?oU%_cZWV&u3t+Z|FkjXHImreoXL3#Rp?hpDqeXM&Grd~Lte>FI(K|DTNmsB zI)O_Bee=$<-NQ^V_JK9aocyXPBW?$l%jG~F#OGy6gDlphl4uIgu|6* zq4SVK%+vlj?9fbrxivai^nxc-rAI`%Yw7LmA!uB*kGaJ^VsrKl6^^_*kDdjo31<55 zq&mkYF3Q;v)Zr&LK4lmkNa<&b+pE!H>JY(Qs}JnGe!lRqBddAM+2&;_3JG zLe#n2PZuW~rn!^Wf>Xp92Abl6GUpsBR7n-egcZOvJ#TulvzhZRPh`bKA6T##F+KG@ z3R`JNw-&F&tMOkkD?6S79JgbPqczyB_b0vUi^=>)F!p$h6Bo^&K#hOLJ7f;pb#NbnGVPrhj>I`9$rFyjJ zb*#YC_bu95=VHFs4B9BSL-=`RB?KCM5qb&M;L(&e_HE?>8a2X$JOmZEsnCe#9Mcxg zC@o{P$*ySeW;2YMb(4iY$q-C-98YamImn&cj$Z>0z=)}-+-?tZT=u{XRXvX3UP~|X z%#js3wb@qpC`JgCN}n@X7cHhMx|=l@ZDc<+ZlSi?S}N=OL}?}pbo})x{3^(0xf43r z+59*`_nW2UIA%4qH3gz$#vzQbc+ULZIN)frD7Mh>9aKpu)5~Hl+_5JQgSUxe$*p20 zr@Wlbr>;Ov|4gd!0^Bxc0Y#iHrKyWXGCkpV%Hiecd3Z2e;=r((@1%Dbsn8HM5?@*F zA=&F5xb9tyFubpq`bwiA0E&FW=R&DEYxo6t$2f^K{4W5=fI&~?Sxc;<-~J5UzI96fK+_*b(r(WV(%7taE% zsUNF*cE!>q)e`(J?MeG$E`y42FM1l)Q-kpW>_2dv7Imbvna`1)7Hu*=zU>t2J@AxS zdkhx@7D+K_X3cJ`*2B~fYS?{i0a?EC#@X9D;AqS~vdz2Aoj2b`Plmn1H`&kVbmMw@ zsV~In*J3P8+D15bkuN;kP=S~F1cL5jS#;1+qpK-DnBeR+a(OkHTxRz;H}6HQxy7{J7pdp*AnJD2VMlU`K%zbkP0W(0YD+(-^JkWD*Z$2| z6@K4*#rG&QdKHe{S57kbhtEJpV+S@q%BG+wC1J1dCCyfg6(}msV~$rw2yV>VNq>!= zK*-xm=+r$!plRbqf%m%D;^z<91dGe`aq9Toebfvt3!xs(z4#oZF9Yr2=V;Ulr8r8nelA z=5)Asu%MP#!KyW(uvj(w5BH9;2{Lb)&*om1x3Gg9Ivv36`LY!* zOrDHoev@GMlzOmyWy#rT`GK9{EvBUI03Ld4;q=}o=7}N4*(x^$eD~EDmz1UQkIo**; ztIMomwdw`8+gON3+YO8}RMChY(AS%tu~> z&DoO3pVIK5_s_QBfeWVW{@4;cq`ManiXGt$t1q&l{a^TNpWN75(LkDOnMK(NlkmRN zaO&EwhOE1R-F*lwZt8w4IgyCliyyMDzSGQ7PWgb5q#rw_!XRe#RyMwNGc4WvnV%zn zh|Z*);`_!LqG?G0zVz$@ON)c7woja`JJHXQ97oY~HG7DZwj(9mli(h2%BC700#Ug= z)n(sA*)H1&B+n*L&P7R@X79uJZS$}W#$xM6Wh(2oWM7LdXkAtq7^JFm&n)Bc!MO+= z)u~Bb*Lm(|{8>2Fdi?tC33I$Ejh07K*rzd(@YMS(1}@aa-lc{t z{7M8of?UqwO&Hz$QAu<1N04}2N_AntY_K(-1H&d>VV~z7#d6+{Pora5u&QN9|~%V(0y#&&qpFbeyXN72Kd=!gIlWclsk6!D9l^Q zvvCiq`AvFf@yYqyTtsdIC(i4k-K^cr?Pd^G?+wL4k@xsjml9}u#dw;~G8&$4oI?uj zO>lW$KK3rNrqr9++`&j$I@Dgy_lyvwwZ$fowQ&e7`LY0<&Ni|4$sd{c#j`9~x(h;r z3;FkVGuTMEm&~Q{2Xh|dhKlkd@wlfVb4lpr9p7JN%bkC-y6;8sS4SW7ofS~wxicGm z*$j@PPJyp$SM%>r%FV|o+&tDqYIDZLtv@ezg zeQt!#I4d^!+boc<6~~-w{oJ;eFIACa^>B0PXu31_2qYG~fOey=O!v@N^G78E>!GBC z{Tc8jvl2z|hM5Xq*x|skHZ;Q8@WWsj_Y3z3)y<lHFM;o>@6U7UdS54W&&7w52bIVr$0r?CFk-Bh^k9Hu;Vro%d&@YPzD zEnhL0D(zay_p%~?nDWFDw+`ZpOOs%ZWD>;L&IgyqDfmXEpXp!kLf-}L1P6Q-Z2xM&-=hcS$$gC(8uaZ-onCPQZRVvVz`oT zg-?&ivF&RPvp<$fm=-)AmTX=?_xL{6+H!!M8(Kt@qBr4zBfY$8Pb2sU9%y^ftgVjdWA2YIX;1-Ja-S4I`?2z;+(Nuc(xs$KEDHASI4ry-+^l9 z90ZLfO?+v62pR(unM-*y{d9_fKK#mU(cc8&mm6Vv{A3pUZWn%;wi&v{O`-8uQ<=Zh zHK@vqrA-oZ*_F$BWYSy3=R2H(;9?Vsad2W6^`67M!QNQXP!D@o zW{frceKV76x0OIeUo}iBbVmP-b$lAm8-9su2G)*Bn}i_f-mB_g0|2Hi)dPX=kgh zedSdb7LsAUCnZND!>zhd{NyakxO)QT*yTal>J#XYW--j{-T~e@lHj#DiAi^fqQT%G zO3lAWl?}HcTD?1H0d-lV&(mEDpY)|DR<_S!vb8dF zP-z6+R@zIw@8jS?^-xTkwE;}aS8*nBV))2w4lVy~NGc;1qnnl$Y>F%a|3QVoj?04I zgMH*3bRU*HJ;xr6-^~O$i{P>USG>5_2-FuUvdfEHFj8~}76j;Ht8qH@&2Og-;XI7I za|ymG+2O9dr`5(~cD&qs-|FP;)2LY1gjwgEVJ&yXz)raqy5|m}nD?9Tc8;0P)Y@xZmG_QH%ys!SnqJl@vN#zxr~ zX0=oc+2iS~Y!%}Yo~U7jjS|_s@8-6`8FnIj2N>H~u~nf{ac->w%Ua~mkC`f9hmOrf z!8~bFonuXYW6qgR{KF9)-Yhb01CZ*TZvXz@sF?U5Qn_R7n7M;me zQ4-7Ppt&8*tOV4silQ@u(IlI@l^?l&KWZ8OWey8n$xXV9U3dYZt5k4h z$z^u*`7ZdXEJX*tm(qw8^(gt_G3%K83%KVpxHGk!S$xr@hVK#JD3S`LIxqRz6}osl zUl&Ij*7DZz=Robw6p*Rf4}}fx%+Fu~B}XaKmhKna*w=ONOjj1qzB#~WPu4L1wPqxy zzZuQnO;aH8%08|;L6kOlm4bHkQ8?$8&Ye4`#?-kvY|q&w=5DRT9Q4!qoY7h+=YN!` zqB%wn-_A1p+t}&#In~Wl#8${FquFFXmazILud?eeB&W^9?68#}&KA+KEBBxY`*^3) z^AO{?o>hA8C9`%@`nD{PnT&k`3;PGr;OWD$M0Fj08}kcBSyX_eM+b9nPU0S73{;v+ z(IlsQ*1hZkpP!M%vj06 z9Nn>G0)TQ)EOW7dCi;Ypy4(x^fkf ze@})gS=g!7%gx?c!e$h-NdS+Qtvb{gufmcqNv2TVJ3E-7(+{V`&zM)_SkZ*~ZFuvDC}@9-r$=u;f$&)l z7&`6-JMnxLPYryiXb~esDtH03S}(C6#5vDKGgXjWo(9UkfFS z-E4tQGZb)R(MoWd9f|Y3%&9Rel!7Pv;M2xcta0}OOy3(0@3ZcjE31qmi{s1Sc~Ufc z-t!ySH9b5~U`i?(TPW|27}{>?WOXmb(d*0stXtC;&-mBit7!_X^JXFUS<{tfi(FxC ztt05fA!!ml=}cMKNo-4bEc}g<5gi?lBd2onHnwIyT~| z-RD^5%(2yfw>6Q|i`gi3RSMg!hoRfRzC>2@0)NQLAFgZvhWEoAX{A90W70=xa>8ZI zNMB2ukcSr|Sxn|TcpY9p7MBkC3H+n+tM_IOSe-bD5>K^p2-pm zj85c-%{`%uLW|p2#rwe?+vtI?fhSoy;fuHcBUySs|`!Sg`dDPhxMX%a-Qo@d% z^de0ZOB|EPFTRsI)h;B2!3a8-q+_0bQ*0#Va6!QMLJ}{m!y~Wu)5Km;jO6Zs z_n!@TO{#!z@nl@g!%=Kf@CR0*ATPY}WeKdBF&$oPPGI})8qxepc`B|xh$D`+V(sa0 zmh<2YJl?1Vqa;_5+4TmNW>^4;hJF|-lTT97U-;kF9zjE^H{od0mvH7wJVv|apx^EgmXvjlG*hN9b43I8 z>U0mAf8Y@Mxkw8(KR887yN03Rn^LM2jpy92$1;P>;b=Q8frWLvGwa!QkNNbgkbCTY zs#aPwu;wqZt!MYKaVrM1M=?M;B_pAt$PP=sJp(@n4t?I7WUKdFV7H&oq_CIH==Go( zg8m#~foAsRlFIIQSlfqwHn+3igRG{h;kMX+n#Q+5^vm}p>xO5?{;=B4Sh zxL=1$_6S9YO~kC9JMqK&@4S+`6islu%+6SBXLY7W*_ZxHRVB}l;$GjGv|rR5Z_krq zDI*SWB~54HxrZw5_ZNa|Kr(desttUbQh_UfhqtfQf&Y{(tT1N|+A=G=TOEiS3!Q1o zxH_62uE$<}okUz^IJD|c;cC=7Aw)V6B#$}`V0;R=(_fn2m3B$d-KO zV0zHJ?}j-|EbX8*!&4|@;y`>s(~fQq8Ah&>*MODn!q83HxK*)^U8s9lJ?L5;PBLC+ zUTIN8*>sgwB~-%LEpK4jwj>;57K_sjze2{QDsUG+1~vnBPHWFY3wM~#U3?3Qvb(XtT-?X1d-e>ZmqpyuWesVee@H$D)Ez~i4gn`+a zS#SB1J)t=I>t1Z!xE#hT_-<~r;Sc1_Eaa9BQO8whjF`B1E;W=!G0$V-WSbR$TFF)1 zkKMBDV5k8^Rc|Cut^rJD$yM9kHm6T}VsU-?Sc+_%3q!({Fy-DXQm;&4lO`TQna3gI zTYZtj+#L984`T37bO^n1c80A8p8sfpSyF3xQ}2r?vZ$DLl@?R; z-K+4v@&t|il#WU(zC)qpVf^f4OfUOWvFv1sxp32aaCkKugQxZ4{=riCV_rTt@=7uq zol3^&%}RJu_c61-OSEh8OPJ@nfSyRmqSO965UQn&qq_3=3z0*qZ*VS`o}SKvV!Zf> z_d9Vy_8=e;6oeFn9vUe zIXZMM{&U z;<-2$3Xi{z#W_Xw%*lQnZ4%Rl2{VUN*aan4tDJ!Kv%{EKS1}m3yFl5;n{0o814X_( zfu=vK*fsaN)c#xpM$A*h49i*Mep-)p?pL8>?WXL@u@O{1R0i*?AAy}UTjAlb3Fu#2 zK%)BHWY-~!0Xr9vOQshJZ;S&R6N9b$+$djcCTb@AV*Q)mRe$vJGmlbbC5$vKvNo=RamW#-E{Dp%|V{QG`s-Ziv*HgdHFD!A$gF?H)tv zm+U^;R-=X0V{Nz@UtP_|ymja8vloKx zEJ^lyUK7;JRzSI5yTK!|4K~CwkY6GKYKOd`>C`3m!=f7We{N&GWula?m5=`_~vT`O)t4^S$&Zr(V|53zpaC^Hv5?RI5p~f zAp&;iKQjx(8WOxcgoZlf@XLLB{3P=TMtHlDvHAIdoSS`+c1c8d$>J1kbl$}>EFEx> z!vhv=af5>Dbm`*M1p|I^EJO%OQFTTbuF+h^^tKe>9u1g-NDIQljrS^I0m*Eb;h7cGf3~MVqu@;F>es*zG%E;>}9%F3y8b!Ex2QObzHx zNH}}A>=N+fU0B_oVXW}wY!vJeheyBe@LHB4^fdcDOiYranx}TOx*-nTo*2;czdW4p zR>a$rde{aNUnY^)17CjX^S)n3Qg3__=OO;cqdn0}8ex=)cu zSC>8@C!QS}*JAqM*8nzR6-Q_C&#<`?^>EzoGIv>71GvvJ<~Fzj{_Z=^a?4EE$GV|7 z&N`Y*99>}7rF7{?+hl{ z$GwoJ(O((Rtwd%in#In%Q&K(+ik}De3Xs+|^0mo8Rpr$#FcGqTcTK0=soWKaq zbr!Rsd)xS9IqB$b*<8H~q;N&TTX>SO3!5zcVA0!`XgBR9e7CsIM_Ndt%&$ybcCQ~4 zOAk_F6H$RyCR;fFCjU#WoZEA7Cx*$0V6NE&Ec;Pm{`##MF4)@)7K+AHR(KWyY8SAk zMb`X*8FKK^0a?@v5!!NjK8}31lbV~8XSKGvUxeae$4GFAb7j+Ic#NOhX~LlON% z{PD-3t-S8U5!Eq4nfTINk`fi((!yOFE|^=ytoXUGxi^%}6AI{Sk`YZ^l>&QjC9#}e zSS+eAY8>HrsDpffu%%;LjMQ(SyK)*Ha#B74k0-bev{gmyMU z`VyBnI*d9D?{J^mf=KLW8|8i7MKZE(j5d#>1f6a4+;=i2&+^1<=Xv~>!fANdUJP#y zdkK?0;|C_bhOMnv8599IHAL*G<-812v zZwIr9_{=hY_f(6^MT6JXSlm*(n&~^3;CnrBTpnBikI%YeOT82|49cYz@ltqev49?g zyoI?DS&&f9p+!R`n>^SX4D1~-;ebD-daY%O`-0FT-5=iXIe~``yyd5^sA6?eRg~#l z3|il{fkl!2U$me^CQM{t>ZW;oGdmdBRm?_jAir@(V@0=`#O$EllUpjUGc zxb3Z=zv%&ZCQpXEQ$2C%BstWyR>A1fqcqvt#k^pXAA7M%8}FwMcn9a}K+$s-l$*@} zxo<<6=cWu;`FJ{--ekmF=*=*v?qLS;FoS!>w zS2{xZN7m7<@dLRZzei%YyDOIT>Txz{8mLiU#OXC~RC0VgUWf~0Ju-3wnA8ankLLMv zKey9r^-8Y!;B1JoY61JjQV?@)326MZ1Fx+`tm3ybAO6XaFVsE@i+ZoXm5{wGAoMc? z*nVR7~!LG|iu#VjTw)l^)W=IjUbRqd= z6z3zen9gPXfuFrk*^Z0q7@(d8zY^Q|zj~8MBRdRk%@bkWu|J@0!6Ej&F2;PG*989Y zl-10dvUjsqk7)jNcp|w5 z_4A3%Hf&JEI7*DZ&;AAursX@2qxvv$Y#5~sIY$mL@R4U#(=2JutWbI`lgazt7s2!) z+U(wqk^G~A>0r;8X#+T=C@)PfYJ&81S_xm!p zdx@<3S0WsBk-=k^m($|T3b0<#!s3To(5u=Gx*5BN{SA}CiIMARL#vPiT^G`-%e(Mz z^f&ZPQ)0O>b*yT7I!*t&lm6_RN?|o&WHQnLcAk9AmNkwBspu#0-sCu0JTbVqIk$`D zzY(%0?+~0~1~964q$j==n4a7Px56bs!=#7YE<~g8!$?ka#9wUHGQ%mgx7qlF=~VRV z20OaNjr1e$(YgzneBn_wfz!`z#0gsQM$&1z8lj3|8)LEAZ@ir{xZ7J`(eJ%}E6Cq)*82i>O0*@Wi;cAu) zH>^z#n?!(`y)+m%*nvWw^{MgtWi~cWgS;fSu<92l;EIGXv_G3y?KNBljXxa(;c-9i z!axq+2Qz)}|MCW^i(a!U4?nRtX_~xbgD*ev^lK2fE`TG??pC*tKgCp%teI!@Uw(^R z0W)>k$t^3s%{7VT!rZTHEO1C1<>gOEW1c&6w0aPDR;0 zkN6|Y=5v9!JV4>AlHl;L_2?y90i0zm_N2vwxach0w%{h6Xv?D1gAv@i6hoMr+yZ+{ z5Qhy)G2i7ajlRbhkdN$lP@b+p-_1AEg~UYz*wPU$ZuL%dkgsAdmWkn}I|($9w}v7r z@dKFUJ*pUPfkW0?W4BETy<4Hl7w&JOslBQseX5aO5u^&2P9TbK-;zTzyD{AuE(x%4XJFy*>Rp`K|!9`FjpCovq# z+_J!LXFA!vF|KrY!CrPRX%kKM`AiAMZk%<(V)03dV^w8c9TQg-aD&slO z;`vO!rji}J<_moWr#bz}IWTXS8^8Zb1N$Z{gNl+gyziaERi?G@s~4`N;*Xr*jPF5V zaa|D}v*`fIPH`~*GDzTl=OW4#WC=S)&lHxsZYQH_7Ra&^1rI9yg`w?QLNPB-;d(a- z;pY;6;h--w&_dM#bsHsw-%(%qA*d2lXHKKJ4~Amg&N6%|UnrP9Wscy}?L^_}Km|ct z*le;}o+$Kj*dcT%)E5SFR+RrC57Le6g>U05g<%>l6l79K4dc5JWQPdBCrPNLI#sy+ zt{zT2TZf&;TQGFrNA5;eBBnG32tyaYV)oiY1x?c)(^u1x_^sm>w%{6+(!7dSa=Xkk zz14+3SJv};0-muEc_H}rMVcVwg;pr(}6e(Crf4b(;#A=N^4p;KT6nh)BICxiOY-)}#DJ|Zg!+bAi#Q+=9D zcQx{Y*PqEb=O-)PGut zO-??zyM77ExQ}C3Uc6#9Z#cn?n&Z_T7mTTJwI=^F?FCt_3mC`+93>ESIEGFwRj9Q) zAD8B9kxD<$M7{)}$>^mli~okZMXj-M_G)2c?oK?HSO7OqnhMr>)PQSV0os^!F&%#i zX#AxrSY|RBZGvP4{RWdT8s1>Rwlo@|t1C49+|S+*7ZvD=juVzdiU>6WZ(-?+@0dJp zkibV_D!~D0_{P!;YI8%I&n!&@J3=Oe&V92?nxr7A9;+fwk}EgEU-;g7SEQEXCyVA7mHf}P6h)MWFNQdSIQ zU&NLRmsh?dhw+LuWBodTX%(XpLO_G`59oYKMvz>jEKroH!FG+6xcQzAJUnv%*Y22r zQyixWa*cf%wWP2Qvbm@e6oGeDhq3%8!-eXTvxI}}{*c6smHceK^#YH!b9ig9v~b3h ztJuAuieCFH5ZZnT7W&;0VXd}#^kzsu4%^T0lYyv^R!R!DJTIa^sr~qCK@}$MUV&8} zewZ~%PHn`28(5j@Nq*MuJ%4JjN3d}sHE;9^bGhxst3-|t)~NAqxxFhd01D_ zu9GE@=+dV*)>){kJ0D(0MhMip)aj5*C7vw31Wf%K1uY1|={>6jzj|EZ(n}S=3)>L_ zZJi99+tEvMgSzST(fxSfXek=+dW>KEYB7v^IN;Ctp@f#0z_J`+)Z08NytHosXOogn z{Z(IIZ!9A6PyN3W<`hKC|BEp9pDaNeE0h2KjJbd3&l$-6bLHDdRxT17;0^v$o$N*b zIs6wV^S{&h6#o(YTlt?1&Hrir-^VoM-z;Ab&;tLvweF1nY5U*rO5xvZN6h?xZU6Hb zIN3`|PW?}%v4f`lm-;{C%zw)Nb3FgAQ<3C>F8(P4cK>iaMdq6Q^D83q52@41Ug96h P<$%o~k$?36rTc#XiF3P} literal 0 HcmV?d00001 diff --git a/checkpoints/gail-ppo-options-setobs2-Feb15_22-05-38.pt b/checkpoints/gail-ppo-options-setobs2-Feb15_22-05-38.pt new file mode 100644 index 0000000000000000000000000000000000000000..a2bd8b6b5901e0e37e9c1d2e16493220fee34b3b GIT binary patch literal 16087 zcmbum2|QKZ_dhN($vlOOky55axO<&asYpqw6pA9H!IfyBG?1x8N~V+$DoG;Ty>`(k zrJ~ZLG}E9$g9iWG^L(H0|I_pP|Gxj}ctym+wFf_n_?Xne1 z7KR0^TC+NIVYsy{U(D2Wo|F)_M%dDCX#ijRZ-}h(iq(E$v;MA>McD z39l9;9`v^~E51UIl&f_ZU-2&?F?=Q00Wo~#zXhm-@l`$eYX97k^*?u{{)Y%_zDAIg z+j*HVzUJRVtYi6ug?xlWwEmW(9md!3;1BtSh_(4&M0Ed5MBmS0S7RDdx!5{Sx5u3k> znExSV4BtXXL^6zT`InFwzSVCc)_)5a9mcou;E!1w#kc)S$ikHYVSc~0g>Uy)myb1_ zC+utyzI_zm;jbqe9uT@ls409$*8ywhiT?E<_`E1S3Hqx-sEU873Lm2Q+#gj|{BcqI z@qbiV^PQsj6aJ_g&7T;>pY%tS4c|G6Kl!h!wQKy%-Cd=m-xOY*QzpFj@8bTy)qi;r z)j#y#!qVAE@n70MfBwj&U&u$#Ts(4Y7P>y%2TldMK>yJp40UUPdmElWTZ5ON>E}ay z%##J%tv?)-6$f&U*brRd5r%t9En$d_DoxlDM^dhL^7i+~30BFZz=pOA5aYH9yq;z~3?V>e{2kjGy- z5^N`Tq*j|$;FPufcz51-JT~e!ipYdv`M$4cv#J_xRt#X~VFS6-PqxFDoWt1qVmu6q z+={V9rg&31hrT&84Ho5=k_B?^@MG{W+AXr1?ixGHG3VQOx=Di39R+c)|3f;Q((t2R zUK62R{4JasH;o!3i^0>_798O6j{dwbmI$P_)u={S!onrT(C#~QqKwv{F=b{K+5o!$=9EA%nNDg*9@DzKwY73k2JDd?-%4|2R7o=$HbcJbZVdW{KK z*PaWP%2!~IT0Cs>egQt|>u`Q{Jw(_XqJfuPFvI^d`lYKt;FlAa*YlD(R*WMqPqxBK z5yXC)0V{_m;HA;o7&|zNysBJC>=h<~{(U*T`-5=)$7{)xgudG7q}kwia2tK7o=@d} z{1O=VI>6P(%jnc0B6zOz7Iy7U2cxbV)Ppml7cEQZ)~ADE^^;0sYVaIB>j&eShApJY zaxz4==hMqx8eDeZA@p#43v2K3al{Kt5-})?{CeF$4lW%|+tYbmVY?;Re+ATj?g&@6 zjuy0r?ghCraa{NIGViF}KyJ;tM09ym4V-&2=srn-tW)DO^#Qc)5*@=VX)D|9=g{nW8D*1uC{$HHPqZi^dwrTN$3&!WXTZ9o8U9kMZGBg#Op$*^5FnxF~oH{RnwN!;;LrZw$nSf0H zb{t*2gK_(eWHNlu9L&l)Nhe2N66);*GN`T!;|&Hg_lQdfT; z_Eu1sIf(5VeG={(dJ&bG@py02b`VRDgkcA>fd6qlUanTf_Ffs*@4uL~Zd-v*wNAii zWmghlUxjX8B-!_e56B^9U0C$8A32LjST|<}9CB>~>)RZ-T)B#OtU_VtO9B(09>P5z z%o%;uPLz>3M$iRB5D*chnNxlMDAU#$Ik?gk#xz6o*rcCeTHIcU6Y0KPJ)qMsU1 zkj|h2jP6tCzKzL*%}Z@@LDLDiYB~;1UAhS0CdhN=E}TQ2svD-8*#JqF1#XZVeBigy zq`Va{YK#_5Jsbe7H@_2~{_jxhHInh)x}!Pg5An*U(b~Ek585f9xadGgUZO$^uLok` zs08|AdKJ_!6G5ZAU|z}`L+GBt(E=Z77`A#G-4~K59FMBZ?%j6kujR`tJ7osb0_@;~ z$06)Krox#g$Aeed9C#o-6J9@?fEw2a!rj?TaBS`h)XA1anTa{5qW+rBuh#+}k4fC( zl!Y+UY&%ic90=*!wXn-qo+^~O;He}>vk<+&6*+2sU#xR=!4b^wOeWkIZ_ z1{6hXp$h)j9cNCaWY^O+^21#m4b`jZ;Z9Y7$n!m@n63n)3mb9RqejxTrv`5c1_{jW zPozoX9$-fH1KtfASI$gkA2=GR@$lVVata=kGMhY%beRaQW~umXqZlz=b_d6K8$n9l zH&XgThReLxNP;y!kS#AS6X%QZ`1R;;^nJo(ce)tvGv6)nzPS&I61RfG4hixgE|^U| zf0l@adg0lyEj)ZR4Nl&i2^ZGfq&FP)LF|`(sJfIvYT05i`d(gZu)&IL(mQ}>SM^W> zz3I4q(=X~3F&+HWM8N0JUR*zW7BRhi00Nyplk&tuB0IDd@QpNCyVXyq13K6+E1CFT zHO4+q50Fwzr?;whz|U$vMjD%9rL#G4Sg*__xV`{MHw$=r>lmqSk;d(cJk-bvVnaZz&Dh!mVLb>GmxHzE# zI@4mQhqM&PKC*xtTEa1u?t~*vJfNa=GNi|v;qm@=w1a;JREiXc?d=yhdT<#IzOO;gMnM$ykR>_++pPPql@?`>HG=x=x)V)Vc- zIOfz0wCOqohVMLhP6_p}SfLOjLgcCUmu>Ly&NreteE^Ey6+nO$r5b4+U{x>*<{ek0 z&rEaicRwA~@*Q$nXhHuj|5HEl{_H1fXQ#3MxBc|*e*5S9>D-m{igrC6DPh9(OZ*_p zXIseR%?H7C)K;u)r)0&$T%0!Q1Zs`YCC^qR;=Y=r=rh+JLmt)BOUFiG$TeqZ7xiSU za4LvOjNs+v7qIZ`fsilb!quKV3}qTiaP|D3keIcI4V?4=dT+;L>ykKhmcEa*wbkfr ze;4eFG#J@94l<4OFegZsRe9efV;h)Ykn3Ec*K{7wu5;i%O*A2{<=d&rxLfr0&UQL! zxfVS0c{sH|6%^Zgs6=E5SnLnO+KepRIxh`S-j5AQ z{{*X~ZP=h?EpWd|3;VyPaU;h#5c|R^u2=3Si0yblcYjaE_Wj}TP&b@S80COb{#p1@ zRe@_c-oL1An0~CO6{w{WSdAxv1z7PgL?4b9!pMfXD!a(>C2)Tn#nueNd&GuhKrl}=`F>Z zu*Er$=UsS75K*!kb`SGn8@yd)XV)2+!Ip;+v8jCX{iBktP3BrGG3s; z)76l+xta`dY6V*{M{f2C1rpsU1%0nn@$B6^Ui9;Fo`-ZgPC9)Rc?&tTauX+Q8ly1b z*M9u9#}&d?29v!7X7nk}W=C$_CWY;pyp=sGaixkO9XOVc&12Is`i41X%vwhJzuu>9 zx4+T)>^O9IECaGOmE%I=O}uM$U%;x! z2akA;!{nQ)Y_g&Pym%moTV?NJOZPgwsMUz!>*8VFf@lnhcITpR$6)7+I9#Hn$%#D< z#zX@b)T}Y(l7_~hR$2!S1PF;;vg}^53U)Laaei->k@*krkt?B^cz>=Y%3ZyJ-`mbm zoBk$xUAYv)Bx6~it3P`-y^1b)J`nA5c93UBBRKn9`Fn!Ll9 zxcwZdiJwL!thQop{&AeYz8E`-Eot)rf*sO7@Ii7KcmJHKAS`PdYv7EXV$NTp`2t4{z z@qX}e5^>}qnSXLRPjq05AU)3*lzI-}3=2mhXqMxW?9M4;6 zgU9=acs0HdBt8vbc9GI3{%Ijxuob1wmyO8&m6yqdlWmTtZkllBQBmmewF*O(exda4 zo#e(oIr>zx7{#>*V`6D7s(th)qxT#mgL}id&v$iDJf)s`M_z@f_j_^jk&76#T$1x` z$wU2WDLD0gHGEJl#KBFAFtb6Asci^=93LenWm-Xx&yR$H+}kud@)lH3|11kHN?C z0KR%(g6rlcl7k+;T%uY4E}C9K*U5RqtL{}G-~9#jMoK}jUM8J(!;P zT0bB;A|2LxKV%|v^V4mYaYg+BH!g}$;5P#dYvZs&$E%sB@2eh(n=V?C~!-Y967 zY=nllqO4Bb0j4h;iVd=Lc*IbHgSK;EV7irCNb%x%Y_>t}I<{W$DDYYx>$weW1Z}-9_;M>64qiD=W`5*Q zbgeEpXjajr5+$-W=n8!nAcj>VXENPo8@XZiiGtd}Ye3Uh1%mm9DZQ;so(AoQQ{t5n znYs{LKAaR(dOs5wL}lZ#kM}V!xCF0%e-FxjLukn9YPhUE7=M;;2bKI{D%)`m;>>QM zyo)p!SHGXwB^aUH0h-2-LO?=aoX7T%ek0?(cR=5X^B(J2V#-PvP_ zFE4gL;<6ENXSfxcD<`16?lw?#p@C}bP$Ti%<_X%>ZWo>lNo^BlS?$ip)06Hx4*1h)kDac^Y=n*WevN9xnS#WNLRFKNNPEF-od zl!KFxM7Zv2GQ@3KC_1EyaUV{Xqwb~-(4Req^B%9o%4e1FCaT{=maIs7Rocn-{}UY^Xbeo#V_Q z*{ohLpD%#t8{6@iq$N}~Qk1S0MY{Vm?pCwLhWbY8V{?*(kFg@tBNOq+%*S|Pzbn)z z7t^lxEtsEljDFUBhO)kj;4@Qzty6x$1VZWEL$O3MQ3HC9SCIGH-q3_=k3nzp66)8T z4nGp^lV9zAPT0 zCT~&UjB{1k!zD{$&atPsTP}w#m)1??UA^+4)*&P8djaCK+0xN)thpQJ(hRrjUQVq{hoN>f_d=@i=-aV~> zR@;Vw^VZ`qXi6m=S<}yxK9UbwTr^J9xJP`Wwu0ELjTq5WDA;lRF^Sq#Lf(sSBg^mm zlUc{?u)RqUR<~JlqMj?+1nW6uB5yzNGvZ;1SQo8aT?+kB4sGc>;bHxDyuNrW@muPE zYws^Xm9kINc*5nxE5I)lxBj6?a*c-&K-*IVy}YA$ffaP@m)j%#&p-g&&VmTI2%xJ zSriV?NXHI;4GbCd7(>_G$2DfkTuxONS-ZOu6&pX=i`aF8RDA(#9<9c-olA-6%!4?| z+Z_h0HQyX^H zB7%j7Cu4$-J#0Fi3_kv9-1p9X*buc5ur>}0`koMW^8`}KT(a@QHk@gB4bl{f@cHl_ z^7_kQFzQ)M=&S_V?DLd74!FjXf7}42H!SG7Dcey0gE*5w1zfq$6dTqd`aO2W?S&&S zL?#$QOR`D5wgM;jdM%bLmBgz>*FfJS80WFC81-R1cZHXNtxf@;>$eV!!hd4mvE%5t zeg{f#l3`tQhT;@)U)nMv1rwk3VX5?X?i@^lsyFvZz1;~U- zLrHDkgE%tCdM4gax(-vd4#1;0N!B@T3ML$Mg+6Ui>r%qtdy^aDj*T=tpd4$z@-TL?D-lmvL8fj@ z;liT&$vm~aOrS3(nA*D)<8R-=aw$DdCOm_djr|1g8in&}o*nn;>m}$lK1{BLTfrA{ zv{w3f9spV@o|e zJh&HL#GEEZKepr5^F_4u#t+i`(F!jNks%k&r{mIi1v0B!ymml)9^9?S2lF+H(Ii`r zJ97Ckk?)q|%^!m#ftOACnvc+jG3P;9u7{p6Uj_#}XJOo$Dd4^3Br)#X8EgD+dH~c;uU0i z%dm>xll1tHEa(ewq5c;vv0$AIr;{B`L3AgT$uNBza&zZ*QgwksQ`V7SDWzX5-EAy2VB~k-QUcOj?V{mqwwA+84A@oI-!<4#m8+ zwrt+g8b~?SN!>^$%kMD&^OOnP-lbh+#-lKtU#tVA79s3b$Vw>FjUf}9q?u998#4H+ z8a+~)3QvE8uq7t`V1|o$El`5>(pXs z1Ud^eVC#P2UgQW(rm}Vf7{&xMhdWL<_I@}x7j@D6${gNJms1#Ou>x*-KZXtWQ*qfL zAMW9oWcs;7-~8n zV%Oe-)QktnO$~&1zKYzz7pHkK_f?7FTpjNDLR028eJdQ>6v{1&XJ|N82Jfh>=8C*_ zlFiLK@d0(^DLo%Yjb2v61SfS?wCy08xhar0)jL>{$~y3RqR4*8=)-wBMlb=BITKMk zOg3-i+)XmT!z&4=tDbjM&kd$LwY4lqqLY~YctDS6t8lANRY9~*8(cEli46%4aCXHg zxclfd6s*68QC}M{Y4C6^zT6A8Dyl+`)o0ATkRY7Pld(>IAScp$60`T7NB6PYi21@? z-f#^mHom}uTiV}s^PY3$e1}D5FYfap_Pk9-zr!SGNcDqN8?Fpm~7yGFygmlj9p`cxE!^_H#ZeesVxMD;f_+*%1Tn zQLG8{*uagui1XoQ_>!;_*!Z2G{Kg57pSQy=hlhe|g4|(ONPF#aWf^R?HFQ+e%i`(Uml2*6I>jvh$83Dfgh(Q~E=^i4M-y3=IfP~%;k9c2L)dymnJD-j-V zUJtIzKciBDDZ3rs=+0%3t^6y z3ii17QSTF%V6??r{7|UKj?TJ-dq-<=hHj7OL#Icet|>*O6B8k;H(A!KBEyNz8_@FexJnrNA_1>VjR{+*0)wz(p%IL9t4SWyQP_OQ|i*&K3c!6I&g24PF2$qtD@@?ERROE=ol5$1+v-3VKlNDIQ!o z6TE(0ryd{oa8)*`ygSRIxh+fQLi6W^Am$?~-1gdyjX9l^^WO=FCdKe9UrYq^qJyxU zn@k28nBju{DqJS7%V}|iUV?r*wfH_oH2znMVY(^QL+BS8VEdx=r68P_+g+Lp8yD<9X6G zV-r}^oB|CyJJx+)c^PR7Hw8TcVF1k6U9!*Dkbe9KQolLM~=vU#<%@M9j;%RB`$ zChmsJqoQ29DFpd*HPwrkihv$cM!nZp;_(XI)N!y-A7Pd-3 z9DfE`veAb)jNC(7zUXrgHpxTj_&jVJF_?J`y-Pw(eDKw}>-gkk0Zkv;hB^CBK(*(7 z`ff4DJ(cgnegjqXa{CS;?lzs^8Y;GQ5A0 zM;caS>RvfGJ~xx(f7=J)-+^h;IGm;4PXm6fW+Pu2ar&w?Xy>8>pH7<5uVgzn|5-Y{ z7*)Y-*?kn86XbAcSt7laoQN6+tXW1x4Q$ zHlY;vPcq>0PlzxRz+2+~@N%QIbZ+T1v}#(A;k zhd1cYq)E8A;TUKhR)UTY1f4|;U8j`eD9q^1I*wmw9kaxwN)_+DYrVh5LZEQF2g za`C&U@XFTL&`h)*1o;I>H)r6f6^39p)d?lCWZ1<-V=g@ZBAOoE$OgZ7N6uB|(5#|Y zNN2`k`o*nazwjg`E0p2trDyT*4{cDouM2BJWQG2obU2?g82paQa?a!m!X`$(S4=>& zr0JlsLO7XTTmlQD%28Pur(r7G+tqtZd3sN^LCnJYgLhLcY3m!$js7s1O|dN#uxpdKu+(9g zen^46^i1G-%U018+jicxFQqWLN1Wvk3PA9u|)M!5=z#k!Xho71PL~zMl|p ze4SP;k>}*DC_M@MS6jZT2GEyMBZ4>NK=36RbB%)5Uu7G=L|$HLXKuzvVO=(sY6 zO);{<#O9|cUAhMs?tX^ZyTdW-M;-3ARApmUrEt-5?KF2v6RsBeB1%0zl4&OmNKVsj z94K`CWZh6lYyvK0i8B!Ax6!6hJ(uZ zc-iG8=(!ZY^w{~(TH^{D;^!f#t&DuSFpiEG8;xh`pTgIKVlYh67i`?oOaq<)$d5Y& z{m0@lMr2zy4Z3=jbKz9VHI|C z4PS*?HGRXutB;YxGD^5;{UR1rGnC8yR*r>IZ%L1LFSU|*KrYycqb)av%X_{bE7tD7 zuBdlJut|!&9csf2ZGY0$n$=KrcO}R+>2uccx5=z~Q)%>*wWzDOhPDjz!WCb{*kIWh ztTgEntZm%sIQ98C$ITzUQ%F^Yu1DsS-!l)g<14YHaU|sAieb;-T5@5~epFA*rcKk< zWB%(#RFXeVzZa~)(>-1|74PAkot3pa!(8zC2}X5(<)Uibcy6)ldKz5TN?Ri?!dR1| z^vo6|_~y4AJMu!nYilf4f1XOWo|wSyEmfhmrL8q?yA@%^MHz0;shu$PMF7Q)1mNqqNlAroB<)BONL4G4_hlErWBf7U+K~qvOAxRTjhISJk? z#9~Wbx1g%a4{v##CJrYnfH*CHaqA}0ngiRppk)WSd*m_ciz z95+!a zuN(*Njj!m~ke~R*s~zJvjps(J2do)i1nD(in4k6r_MI9}3*TLYR}uL*WxqHlc{78z zGA$K%*YeQJ@*2VV+ZZEj4<@_%(QM2w-nfrmu-@kzTsAmD)pPuK167*fM}QcF@$w=1 z**Bi}@Brkk@nQ;m#sV)1g)Ry-pB%pdHHkwQPi+qR3;#ujiE%NKOsHFtcS94h#}tsY-0bP2#} zqsZs|>&djsj^HxD7FIjXhn9>P0_ly@u>EWUT}k-b&%WPndG*Q-B-Iu}%PYptfDuiQ&Pud^?C z@x`69&#J_odkP_A*hi?Erh%)LPv9)d49VMVXP{%(6+G{KkUJ6_i?Gd>Z7H?|wObdt zJ)ycd&Q+iLai*UJyC~u@ZyQwc90i65?#*3LZDwR+QKn?{#duUqHs@@55rl zgWz*q8w_J7z~hKkJhgcMT+wNxbscHYHhnlPT#`bUm|WrJ3+Ko@@MoEg8t4}9$GQF%n{ z{EV^sW-m4Q)C@N?j>D}aOK#`@?b-`>7sHo}5-@zZCQg5y3dUwOY!BZPgQv_8G(F$J z9D+u1cU{z(?(+mVwowcZe@x|atzJ=^Rk_@HQA18FWe;AC{w(M+w#A3;R!n}zYTyqN zh3ZwqI4@CI?(@hp@Vlf0A?M^-*_9h`m_^{5#hqw;y90MbLvp86II1N3_GP!;I z1`g#eSGg67v*5xwF?|1YJnncd&Ti=$uu$~?m}ah6Q&uB@NM$juFI5x@>!StpwzgoR z?sL#LpNWCXZJA_SBrY@4=aNpA5byL+GHc;vV^|dVYvRa0JMEN=!MNEATqmyXVokM17~PL?uiJN zpgbHKhD&pS`%~bptrr|L+rj0?sUd$#JUrSc$}Z>#<3Bshna+aS_*3MDATvV+=gYq* zau1Ra{509wGuvSA?Ob@x>5$DUl-Uk9O>V=Nb@)Y`k9kZ2mcJJ6p)XLOa+h@3gQA=0 zb*JyxcE2J>2F==w1$*$dGER@WlNc$Q=E zfwG1mNTi6{S+vjL!RG?bq*aH>K70aKqLx5>t_*qL*Mj1z3*eO9O?YNW>GGbH+^ojK zuwlb}I<4>nD(sNs%688s?#t4dS;Z`PJ90fI?`KF(oi>#7w`c?PqGjx-1cf0Eck$*~Z_utffbTw@=Xr{? zV{hSI$Tq%$NwP!mS#Kg#sAQ5Nt zYvP#V;jm;;9DM&4L$oZ0kfe<-Nrlobd|iEjbOxje{9~@thV5o#z8Htw>*vts1Sg*7 z>#6WTKaZAlZlWAULR!a2qjM%Jg{W!PjCA~jGiItt4g}%-&;LcoKER+z3;H#>f?4Wbl zG;{)J?@EP@`^Uqqb%$AbSdqiIww0XrzH#vVs1)dr(ZlH@Rq<&9V1vPVn4&IE?~Mt9 zMOp4#Z`V9h=eHShqt5}A5aEm-9|1MnM}Ut^xQl5Wc%bRKFt&X*n`fQF9!@!ki;E(- z3)A*;(>v0*lXD{IdEMp0_2m#;kuPR*iiN(bihF`$omjfluNbe!Il`kt0)yudWNO1I zadnI#lbUdxUK}w2J7j{mr=Jm;eQ&|~DFtw7_5nP6ESvO4ltBGkG49*ab;ND!KHk+S zLpa@wiJWU<5)*Y>$>MfZ0{3wmewfLz(Ahf7Xn#FX&n#rq-Evt^=Li^Q*GZaYT?H@u zY&Nn~#7h%m;qeG(7B*X!@--7Nva*T|-IdBdrrg7=Hw>8EiBhoTf5D@5 zpKGUhT%kqxKOs!@g5Yg_>^PC<@>U;&@d97&xkf!sTCUC{_!m&Ne;RJHKnO870d8fQ zpzM_cXIFM&{f}~Zr%?>Ez2$^1^h~O>Zvm)!bYjd>Nyty}K+SoItZw4~@>R|h<7fx& z8s3k~?&QI>HLp-}`C<&`G5Bs{2@Q&}oLbIecE|eyM4nORP6^{!rUxr>_mnho{eyY1 zPDzU88Fx~r*`B8TQ)gD>Uy*hL!$e+}FZbvS)01ZTu@4urrQ? zHF|v**PVx<3SE$)d%yN(qaTJ>jAfyDx%h5c40$|f1=+V$k_4D6f(0eA@U+_u=CbQ> zz4JPTEPRTmUg@z3md()j*_e}jri%A_6q!}}0rdE;k21g7N!Nf}2qSM{x3`~Q_mv+o zeAXq%vHu7^H$0*rC6;sPU3V~c;AQ%4Upa(UXL2GwhiGK}3%KQwK}_SPW5e?8Y^&q} zsBudMt2>)vl8Y!a74Jv4vHhfa{8FxJyC*xT)gXA=BE^kZas!NZzlWJA;oRx3Tk-sd zE~HT_Ve-er#8XC`y(@Res{IP6CG8A(F2T&T-Hwa=)C$ET?67poFZ3TE!|f=M;*#u{)HNXDZT{USpWHWdeH>(M>zG z$FZrWw7AyqpK-P8PQveR5c->qIL)%VP$LO#j+kT>fr2b}OAFgEW)KM{3D}9kMxB_Xlv1L_^oo9V~d~E;3H( zvcSJ}25YNVWMQX_*yhkGfz{z%I6~BjxrE1a^?|wg(72hLt(yau1F|`Fr<*knhPyyF zdL$ZzC~@}PZtSl(@3O1rWx`a+-{pT=U;dT9IOhKgJDr_|{XbpEM$7)bqnihodj3uZ{axH$ z2mC($cTVzuW@{?`{^5_~-x{fBW@1&pBtWwf3{m*~3|T?f2W(T2P3OPfU#O|4MRv z6ZzbHx9-^K@43LkZNHn2_wGH*%w+gB{#yzX*t*qg&mPaM`*(Wn^OLag_3-rd^sv~u zb-#oU-|V$M{Eia5lDF+~+wbD-=IgfCM_}uIFJ7rSE(iAS+~X(V;_K;tVCNnWm;Ijm z{JeZ!{PiS!1ZT~16y;g@@ig7Gd-@3d%S2+?9xu23>;9D~;UnzG&(qQM5pfjcows@) z@DW{oTzbEcn4L8r?;_ks+?qe!M$|C>bL zN8_l^wEvJC@N?I(wH6hVds;XvnfH$VIc)!*%`C$=|KFNw4bUl9?k^K)m-Teu}Jb@+Ipg((~b)ZBTSwgP6Qgb znx4xIB_IAipf_Bu!58ZoymBiJY>uwQ>s@{%F(-?iqVSURPQStl`Z&NXdbpVW@zH@= zApvM{4kNF842Z3EJ+o_;5t^P;C5mU`X{=KV39PGt-gBqH-bjj`9MC4Y_F`nojanw{ zZ8f}oTMM*c2DivF1GtTYXs6Z(=O1O$vn!X8`uQxWRIbNGBiEr#oS#Td>BZ@1HgrCSr`Jz{6*~O%yp9Rwc%4H>cV+tXP$Y!7 zD$=leF?jd!Iw>&|B#+ip zMFNtJl1uL_$db(-ko1a=&Mw?TO*;gsW_viDWf4GbyBEShLvhuqw?=eANg+A9N}rl} zs+v?zy9MeiCCT8UOs?AnNqR8n9!*p^MG^`tu_wgUB+2R-GOe2QY}i*ENb;tgdB@>M zZ4fQ9xI@IF7Sm%^swTyq8|aew3#9vqpvf0aBkB&i)VWNVe6OvClXKH)gX%;0bfyS} zb#u_c(vDcQFQY@Z9EhFVb(&)ALL=FB;hd@Xc_njOrbCg&{4-zQ$=+j1e&XgkO%%V*P`5PsTNa*Fn>Sc6NdKVYqS zE}cKLgwU2@6zbEYA#2x>8ENOp=;1V|I5CRmMiF3ImII;s=V)){9B5KnNXM3^QlHxi zbVF@CRZ&y`oA^!Ccc%>TUMm4-R==-wy|RH4P9uq#$1=sY8*uJp36k-=k3QfpXNsox z(5UK2axOWTwySH?Z9UFpuW$*to^~_&FiuGR+SOFcVJ$i4!qT2EKt3nk;hxI+4E*n1 z&{TXP={poj>p3%s=K9~nTXGP{Q#CmKaTfXLcAu_(?SyVUagZUHPQ+)KLA?kiuxcAU z-#LeV=T0`sJ+qv4wYRc+JnFGjH;$6gB-|c#o&DK%k+`l)#{7If611WTJDw*pk;}U2 zpUvsWSls1yzcxkFmM9W5sD;57W047b4x;VuB%xEA`o-)ee2*?+hJF?4pO6Oe?M1{u$f#8_2A7r*o%;+>I@BeNELsxO)xo^T6y35!+*L?uF)@3gA0 z2yxW>W6OHirV#ZG4SK{*)}&bFG@WarVG{kJl^NRLOq?`B>EYI1PD?3^k4uYCFtU~T z@^}py2r;6DQx!;1O+MFukq6cJ#gav_QwU@)Bm;dl9Dh4Cdd9&50(H7zezz^T8N8Q1 z?+YV}`&}S!&=j{kT~ zjFqw&b??hmGgpRAn01?5cEf`7HdSzGhY0P6^CbOyn(;-~a`uB8LzZvYNQLA^@L}{; zqFF0Ta{1E0nqQtSE%2Zkzc$f&*iYm2OK5iiPQTMGtJz<)>f!?qEi|H5m+DOzuZj!EE1H64JW_inYDq zZg@BiQ}v|ocCX>%0a<$Jo*?eNcZBFKtiU(Xlc+)X2>0D^Irf&Gr?#C_(5~kMP3oFV z)tsKg7A-C694JSUMzh)Kk`pxkhYuNVn8&V^)F2njlF+I#jrgp11&z-xnN)S3!1LTB zsvD<|a+ds6b1dhOpJ_kP-qRfK%ijal?=9qoRRg9qsgg~Zwq%MjOV=DuXJ5(blF3KH z=>*OgY9&_C(OYuFK~@9gSM8u86UAux=>ufTZWAJZPo1o^o=tc5OOiUBy>zdD6@}-A z>2V^7c(I>#X%D1UYNo{Z%S2pw#lnQ|M?Vc;Wk?pg3ebYKLhM}k7hY_O2F;po99^M~ zBYyHE>!w&$ZNdrm{gboQzW+45SZzk5Pfj9c`#H4Pc_XPaY=Ml+#^jgFG&)5fh>;XX zrccv^iHDoENmyn-?0KC|10<7ZsZli4Qwy3YK8U%+7s&Uye5A|848I<@OZ}gnBH0U1 zn&e-#r=~qWxkWAI5T_zZ{SHKu?$pJ!->HB&h}e^@@BWa|9xgd`$e!?vTUVveMJ9ex z4Jlcq05e~?FbcZTuyUjoGhU~_KvyO-!!F~}FCUmh0bg1@=O!&v*+L%~MqqjEJN!Fx zo~-c+C063GbZev@x+TX!Rc1U5)jmmL@9v;|no3oxytPnlc@;T5JBYfBZNa=9-Edc@ z8FCIyBA&r}$h@8`TCrvg(TUHXYwM<)Y}qbItpk40d$Vj&J4TS%Cmw<^9)YauUR}ax zbPwD683c)FcbNXlNtp=-`d>egc|A*$s7YBbma>9v|=*3wi)7v-nns^iJ< zDQ7jr%tokNGQ`~wSAyxvV&wUxJ9x%22i<#9&}CLQoOM+r+VZ{-{An>vYi>fxi8r{L zSDMgzr)FlgV<%g3PzhZWq_9wJE^IgKWxMQVfgE2g8kN3hF6!=q#HGd1?x#+2CW*2l zM;74TjVunl;e9{uNMKg@VK%AbGY+ph1;Z_$@e58MQ{)}V?dB}z`1(}vZc`*jWCihH zv^y3*&4iwv)9L5XOfcBz44X^N;*I8oth&ks(rcfE3y)Q@{yQ?*r59S^q=gH+?W_>( zy}>Y7K7EGi5r#Zp?-di6)&_B`{)=n%->0J{z#nc?y@$KSzaa`t0U0KX%)JYEGm3NlwF!>u|a`hvzRnhr>#7jL@D! z@Ei+=)!S5nwP0vP@B+GC?kbo+od?%`YtovxFWCMa(q!3@m+*1#L&%%{2`p#Wz{!nQ z@Di`xav%Q1q1{}Jkm`os;V|}>ZxZ}6FUJJc>#Wzb257pYgBLRUaogs%>~1@RagTWJ zr^1b(&3i`@BB`+GeHOl5ig4n`BaXqTYpB83zmAmAJW*hHfcl-Ob?}=J2v4*KS6fk zBw-R)GZS2%2ST9aC^t3M4rbX`gZedfh&U;W+?t)JcFzneWv|e)7HP0nT$=8rN$uBIyWlwVPjPx9`G|rjX&=-ckY9$cjlxe0~B-+a$R@m#qfmK&o zqpyynDHhr?5+USL5U0S5pG2+x#qfKMvIF^7u!i-e7gnqX zKu$=%^g0m`&Zur@js$Wi*zX#oX%pcE4tC8NuNx#`VAGk9)i*6EmqO03e!E(Fr+CA zPk=Jb$kqWnn^;Wsc!lH6!psD<1UAP!hn@G`oP-*FLE%?27&>h}HPEg?^Wa|?KRSu2 zJI4i?rdvqdm1xnin=J85V)we-fY)|{WXXf|*s|s%GqA4(m;4AuVgA|d1-m+^6~Dr4 zyF4HMq=nJg!uuF>P?So235AscZ@3bhUFn0q^H8%yoNj3>#jnX7IDSWs&e*F2`7Qh5 zdGmCz))OQ(?$eoPCnd@62kJ2U)DPS}SDSX#tzxY0Z(wxdMXWCS#=f>A^wqTy5Nzs% z8+ua2Oeuu+9r(tMC8oi7uV^xzPYN?TN@!ULuMgJ8VB5EHaG0ZyC*9o=I8I!}Tl)u4k<=)1ck@)dEcqP@PDaxAZ%=^y_bE)Ir+~@*n0IH-;{ojR%tqiI;LJF8)4;N_@Bj&bNvnAyJ0j{P-CoI+UnU z^hvmJt_APU)#c4wVeq0(jEFjRq3MigP-iutJZL__%$3UIQeyz$4Xx1MXo2kW5q#v{ zjvwFXQ1b_}^vm=#K;d-mtJX0z@#UlCzlWja{!OMWEd_sDHGwSE!m2hU^1%2G3PhGd zd)j>{+7`+wDvN-&t|r*6pNd24oaiM%drscw5X`E(h4=4kgW|>OkTg%5RNbrwADRdU zCq}S|e}stVPkSu0m86#b#&qn_ZPsg4o|MfD#lDpLOxNz85H))O{UNY|6-e!1@?KFq zvauB#VujEuaWdI9e<4-KnGfc-UVwX9J7`A51BBLc)4n#MR9rkRicDmcgGbr$Xkk*H zJHQT!s1UWy@?@kelJS{n0#OA$IL=!SROr5kn_nJdk$N^`Jz0t>+!i2HCFSTbPY*cO zqC(3}6WFTpQxLx61Y;G}ivGS2VEC^ct`ZlZkA9v+&W&xP&{dtbj;w@TFD1x|rAC#) zarZDp-W07QDQwYe#fOa?MrG+TP!~!9YY$Tr)?-Tbl{!!+ydNvdwJ6`TEs($OGAlvT zKusxZVX=emQ}h|8*78T5aLcKnc6QI+PkTD$~W03xHpICw&li9xT_br*( z(h2U-%+zFQ`cuLV&28pky<#_>Q#@xZs1=We8%;PTqGQmMsbu>+{=ly+C78N91`jCT z#Ur22L&#McC5Y)D!j)?lROfa8y}$MboEEglpz2^6 zyoryPPyNbN{Q3?l`47SI??N(oxu2EFc477kiDGfF2JO-R4#hc}iQ+aLGOa@uTGTeO zJ<7eXB*>Up2)98{yc;|DEFbMP4Tt`sqj16QGko~`0^+qZVd$Y4`84E=N0#-$ORvu$ zeXETlAhm+nCpgo}+o~{Ur3&r;HIK;~eF0C3wy+{1u{1hW0S)5TV5CDDSOw^?<)aZ0 z{aX{3Y!rqruXP~MR*V-8^OLh4rktZ%(qt^+8^{_epg{Okdf|O1B=lZHai?>rV<$o^ zhHv9cLl1lqH-=m0^g~AH4>;lcmR+N-4twWk7zY&VnD$+wi#mdo)U^#;Txk_C-JtzBgaPaQA+LwGz)!x-$(!>-sTPFbLTq zWN!{S;IyR|;Z}t$?A4x312+0`vTM5`YU^$0MeZl$H&!L{gI}PubrhQ>B1PbZE-jk4 z6ltWoOr?T{w*+?7;3emcl&?gke(A-FBaMVHVeob9@C zP&_z+HcgF&&G(QQo1=uwm!_e{;R&!kq#n$ZuA))47Q1oPEYNeS#rjM>xVuAwJm-4_ z*4=jWNo$<(<$`&{Df=YR%#Vibx$#^jFAbWy)r$M{lsWhIxeqX@^dP>Ryc^af$HRdK zlUS$KHe8DqK00)~#P7T5i>90cthgRc(BU^o2` zYpXjB7n(o9E^vk9b@F(wE&^UPDPv}UG;Py=iHhrsz_g+aBt!#H_Jaxu>-B+FrN59e z?Jaz$ngi$WN)c?Ik1}I-n9Ie_A+LQYjwB>vdUGAqZL*yFY%fBisc&(8s0=;QyM%Lj z)f@J~ovmo~O$sIARY=>4=eSw?GTXt~3l%0|WWn2fj4f3mbMBsCIW2NDJFFI653GR8 zftplmYa65Z@ho%MkuuX;OYwB)qP}uM2Xad1^HZ)*k@&aWwszx(~doUQ}K@CQJ@Y?Pt$43zKtNLCh7uV8;K* zY1rqK40p9O8TJStbQ%PL&pa7C@_h+i(EA*p2j#$(c>;`2coT-Qg~+uOr_0x?(iPm( zIICU~R_alfZ7yJTxAOe^`$drT=_lsD8-N2mAL{PO$r!!;@e;8!>cpd<8oulhr6w1Q zVW0XX)GeQiS#SwG&QGU-8aZr$-Y2v=X$rZ9#;juJWHJ)0!wI8>;3GJhnhDl1yHcV6 zDvY74P@AZ2(IyE(br@vb4I8fd!A;Y0raXHESM_lu{^rF>3QBH4g>WncRwxoT!%4Wm zLq}3@Z+`>pKWfO_ znjMcDmp+7b1`FZU{W4TOxg5*I<1oGNDaOs4&YljONhSxRpqJ%lGH`Jfw#-=oEF73R)*zK4^Vmh3m^RYe|-mV90 zI%nd+-VW4s900$eO{+i|7JL|VWxS2TJ=5rrh4O1a{ z`&+0#Hl3FEh9KdSAvP^hpy+%BC4Y8uR?8%Fl9_jqqjP~#*N?z?>sG;N(@YZVI#{uz zMFX#y6KYv>4`&EIheM)?(5{lfL_f|3^`-YvO^KfjCscAG$QfeCsdlzx-23KZ4#DYHQk|aKdI8@ns?$;rd&bbn_k?vz~( zBezbFF*cdYQYj)?S;kfhuO|uJEF-qNhFV@(h)hu+8_b&*&rS+pI$M?KU*i+>C|8X9 z2wY9nTXzys-FVt%sDc+ezpy>=&xwb5HfZiX3|qn@VBLs1={FLjcW(5sJslcMyG|IE zSq@>AZ#+F{>rR?GebC460F5|b4?+=@sNs79l&4OlhpKiOuhUx$?}r9?<6=2fMJ}ax z*2h23MGphBdR!leB>h}U3cI3m9OKQFP1YCB! zhaUC@b1%@ExGQQybZrI1h=!m|Hy0=7T*M*o>o{}6AJ(eW54}bsLALD#m}Y*!?wlY9 z`YK9p?vNpyHwDA+rH}0I>Imq6QV!`+g%Iz?q0F3U`tZv-_Rs+#x@;N;OXBrVV&^;> zDs~Eo-QF>KJSoZQT1oT1{l#O(Gs%T%k753IZ&srspDW+00j8HYprRT=Ctx*O=)Ih- z_x9wPO&!6==8N>ql5Xy^BVCM$q8Z6(Vt8vtFZl9Vn#6h=kioo(Y``foV%DwB*@G`Z zr+xujw(k-h8hnb61s3D@t5meISxg!(9HFf@^JrbNF`2ae6?{A<1qmTec(33LM?($a ztY$MA^azE&{grIlZ~+!6YopPLr8r#s8?(D5={ZLo<^rQls+}B3eRd$Mt9yv!;SV8q zw=mmvbd;^i&4BG*jog?&DfE?QT1) z92cY|N-DJV`z3b!gx_%A?KC^q{t9(ut0%X*nfw(T-f)&ZZK6)MC^mqP{{v2t*?dUP z&cOM?rKp}ig=~1-hDKLAaMlwmT&T4NB9k|>X9Z$W*tiq)f@(n!a#7!T3`C{R;p%ZO zxNNfwUtSa=8D;Cy(OCmmea>NO-tqS3rq0B(Gp6G*>$CXGtPfD4409{f5r1dk$Yy}K zh4(NI9y3a${&~y^_ zau{u#KA`!a4JuxJhjlh>!06s(1H~qQM57TSe%=lpBp(=GosbP5#QNE~`fTXdk)c(m z(m`A8E4Vl&!tQ;Yn0>4%3a0;x#r z8{AP`%36kLz~L#*=s)XK<=byJ;d)96Jf6o}0}TYjU4v<)F3JtmFJ6KAG$ndO=o<1r ziNn3y3-HE_bR2Nh#muD#=wyQe)?j-+=fwUAY{tDQl&?e-%TC?LOVJNNu$Z4}4F(YZ z!-Mc5Za>B@5~DXo>hR%}Vy47WiyUj8Nh4nTLJ_bh@Y#%Xc|!Prq}X4_&Xmdz8TyW5?xcvUu;(Zt7HuG7UhZ^^}q-|p-_-ut_wP{+y@ zEvI#X^U;4GHy2+uj;q5(FP)<=%NyYt#&4XF9pe@B_h=0 zgC$L`m!btvR}*C)adPLcEKQV7paH#2aPv+o_dHtzp9Q2TVJ4FRsaS9~dkJx0^htJ` z4ERqyh0Y&KnNaa+cBn<3-27vN=7*(e<<1wZ&TK!Jr{BhU8?Qk1tHJEAD{rtsGy~-V zZsI&$Yq+0vfb(pM5u=&(o%Pz70-x@9!kh(V81B0j1MF(sa7=<}7mh-Djw`Dz{1Res~wIzA%oe0Ec z0DtA$(aAo3q$A}viGSlrgPX-Er?UxqCQFmu(~pA7sSdm`v!CJHWJK#;^7cCwwdo<1 zWpwJ30SIc}gLyDWS3F%}+_Ny8D}C3TINQ#k!n|+c;|I4Hg~0;Idm>A2?e@XYpQ>cV z8ACEl)eq|VGwHgj3M>c==eAA~r&bn^@q7ARd?jB(^QuxoUci;47@LxdWL@ba#9B|0qw=;LIAFC761_`sfl)t`q|Eax?~XBQH>2SB zob#2gE_GORbRvoUT!ACKGIYwl*ZBI43>o|O3)~KH@kgsV-P1CQZq>A-DP!v}w_g%g zcdLScViRjP$DD5Q{LHmJo5rTxJ3*>@d?}gpjG6G)kGvES#bu=zP&;K3x&{9N4+B}W zxav=*X1&J4%DfosYeSsGi|0(3xS!rFd5wZ$3sF@j8uY|#*?X%-@QU|yICb(F*aeQl z(w)B07$}DFi@w2~?j9^SYDGUg)S~Jnb7qR>LMpIdi-;Mh)54SC^wHOC!u|Kt3a+Lj&#epr(vx$kNI!y>trMqM=q4_KxCgDO1 zIq6$O16^uZfeDr6f#P?rqk4#<;1#AhPVqm5ON@fjlQErM{gQZqtbkK)ap8pB@RM9ez_7PK4CpYIOPcaGOFNLmJ0 zUKGKqm8c8Ns@$1g7j$U~X&O!0UakFf#lDjF@iLp!E`bUw;ln zO3%W1t#k1G&R!66%L2lSYg+VDcD92>FMkVZ!Vh97t86b5=ct5+z;Y|7Qk%7`w*oDZdE8&ywKh zIw@M}_?EdE(T3BTBw5K~Z8mVV-aW#%?fm4~##&tAc7qM>;rY*o;h?5_mE&pf2|N77$fX~G zRLVk#%vQ)l{=Le$Lc|hBn`Xh#CM{U>*&9O6_94?G!d3cQ!xlWb!0otOjK<9pBrZ&c zMi)+@ZyKt>S?n<;pVh&+D;wBbN7l2{u@fFQhM;=r44S8*i}!{Xqi%yAV|v+uRwi}8 zlbvGpZ&xwT@9|=Wq95Y^Ezfb|4iS*BNrCZ9VRGkz0X`dg4Fx|Jvvnd4xHRAoQ{KSf zv?W)uptT7?zI4H+fUivccn;i1*@+v)HA%9FE||QO0^_R+)I>3aIQqSYjnQZ5Og#x| zcgWIx8A`-F$pQtV&e2UjUC`IbfwmiO$LFc@X?)mn9(L*^d@TP8 zIsL-eu=_8R%{~XZXVbB|$prmWHqhOi6Y$DHhH0{nr7s`d!Yt9f^oeUXPIT#oMHXGK zG4VFDsZ@c+7EL74)sskZ!Yy)ke;cYUDdb@n(wL8esqD^sVr0#;+j!DT2xhde!P#O3 zkgwZAJe+bV)9+75oZ`sWf>_2M~uqLV?4 zWxsHhUz|Q19p~Z~OMdcryAw6w;jP!J>JrbVFJNl15m7f3A`Nezf^gzuBK&AIyua{< z#sxH>O2jqvtrn(IFPA}RV<-Ef?g5=pyPv3y9U=E8Ehpjw*U(2n4hQU?!CZe?6Qv>t z$hDCqHeKD2C)>g80e<6GJZwpw+1w><<&o@RG=wbOo!C1#jig3e!_2;G5M;i9au3d@ zb9i&(XSaB?ZZU`A_!?Yum>1K(c^N-AjbK2q2CTS~ z0%EidqsSpRpT|cS%PQc1@&;rdbl}M9bC7+k57+oF0?&Pdq+C)JL_=WTS?k3EJp>8cY?%)UhFW;z+d&zjD^-wHd5{pTfMfO z6Pvb-UGUTZq-HVPr!)VuN!vGpZIL*!sjFZHpTuBZlsJtQzXMmEIYYp#zt|2Eo+CVS% z20S>ki!E8a316G{a_ki1VD9!)xcG!LQF#{$3wM8jqR$a{>cD89#_)syfhLKaUE&yn_9u4s1n2965cG&=*;b@Pq3_y$rrV zm&R)Hi-)bTf4PFf@&)Lne3u2jgnb;)_J&D z$yS~ylBHBK-JWXR*QCAEPm%G&iMTvqC0My!BMZgj@vD3>$un5bvDMi@BJ*$a{3kWy zF0h4UYT3}XU%Pqh)^+$RbqaCkRFiDui&RZW+9V;bnSD`INS+NXCpKT#o6xOsIJoB& zoxQ?|bS^zdyk}@u&04mGOdHc+XH<=Vr1mzd;X8*|hjWRyhCH!X$tOReQsIw!CDAN< z40qs)!0)ly~T^+shAQ_U;TCz?X*e~FRez(KOKeT#{3#s+vN-VE&OA`9%}F_UVOd&~WiI&iDC^_J1ev z@C%FS;!nE7XPq%gd7*{lA~Ug4VhcL6bFqDDH_DhCBg6L;=&QdJ{3j^VN$y+eG#bkO zcGHD~Pb#?R>Jv_;&q}&}>=e7Vt%`l$cAR<#jk0OaPvY1CGjiWkx~gxdyvc^7I5PLS zT$Q-$)~b#rX+(%Wg6gL5S1o-sh$5wtBu`F{B(#WBIZaZl`lz>=J{XcWS#+ocJl@5T zQe8XZdvF2uSQiBua^pmz%e1OrN5G`$n_yML_1RUemyXhhCnk|qo1Ss~T-M^{eebD# zl$nW&xd6G?vDjo{l6%#~P1zvvxR)+s-K#WKXR_g(sa2p;%|@@?PA{fs(6U0kDx>Q7 zs$R7U@^oJpXVVCW#HMvoowLC-`)w}yYN19yhKrc|Lo#dg?lK=U;N$z}_AXJ|6VKEf3sa>^nYys^LcEoMMUQPr;@th{Qpw_hn)CN`v1?j s{&mZjz$@aP#IyT{>&R!dX&~-iCnU|$ ztQ0BD(I_-{ZlBNZ`+dHBp4ad9fBm25c~+6Qg$4W0LQr6gCxYXrZ|WUi=}5dZAA zaT8a}#Y-@hD?Zmrz{$)>I*u#hB<(aJjw`u4;I}SPp6XMaqAOFyoF65OcX8 z5z1BhTM@H3t}0(2z9MRW%TW*IYItxp|D}kT@n4E){UK!vSNm^8#{8EeV}C26^S1!q zQ0_Pn?)ZNxV*W2hCj247oI8;(VknfW_gjRB6JHPgzx6QqTZ~~S*T{oA>0f%7|E-7d zA3`j+CVU|xq1?%T35n&J{?^0nZvj(6x#k{Ri%k*SsecJsw$4A)_qVZdr~NhJmU>?N zu@>N3MR2YEx}jnI!6AG@;o3Nj2=Nm9>ppNf5nK}RR|Ve^|FRS=L~xlus!X}lBe*mE zs50Z)MsV%^sG7o^8Nr?PN0m9(K7u>@uc}QUD~#Qp#KafGX**W&`~F?r|6B8y6;S%a z{7oj?+m8H)@$;VqhchSW0kMmaFT$fr$!#R#b0&Go9>&B`&mg&grSlTblU0v|U_*x@ zG+)Z+n0bmbPhG>XKe37ax_$^1&d!4H)*<5TdLLzU?20lEWMq~<9=<%9I zC|PO^TR4+YrmUQ%7D<7Pw+B8cI)eQ!$1v$w894QiVBdWnMXm;f=Q>7&E z*z5t0%$ix4+BO+{^>e|0&lw0_84CxGJO@*yXSmH*7S81?#_HnfsJ*ff_sfljr#t3g zCMOJ1H^+m}VWawhFY-*`KmmF5ng@EZFNoxKZA@S63113NVZGn0df6!@kaO99OcY-S zzIn#TX!*e724m!2dx3$5+cD3FM;68$1c$t{^jFzFsvB*GBJ;+gUtA*Gf@9!#Z7Or2 zWi_>!VgdRIr|?_ycHH)E8N^>VL)rK1@ouv*WC^~tQM|Mabq0;$px6XnV0AVWs-4F` z^*CZQzX|U*4U?BciJ)s)fh}#em^CH~3i8b0n|=lCl{G{6JCe-V{g>&|GlIPBS>-&R zZV?a;G6D6e$#i-6DeQ~RMukHwA=k4Qvd@oVWq--R`cEokm2e~+o>Kq^!pkwhFN0|J zjz!LF7qsdML%Ez|&|7q#e7MyC_m?%nKz;{|b(lzOCXRvQ(`rdgcolLY3&C7832M@f zVYj9zQ9D@%BrTRGOtNC*Zrp^=rb(b_=82q`2v~I?8bgQQ;*6p_kkzh$hEeB<|D$+_ z(mzCoSKWdrDM9wt)k|ca;#EwmodL&m7%;d#4PUrS1^B!jVuvkZXH^0Ykk`<(ViPIL ztV8?lMz}@w8(prSPj@-SV|i>GSbNWbyG_n`_k;z!zi)tFVxsZG=hraINRVBEv+?4# zOR(x<2I=}CgPw&I#DZ&1A9O3?HBL0nt|~8%V(RHu}opInOl4 z0-w!F#+xJhNbJ!P)aaj$!`pd129B(zR<;#`2rIv1x@a`2oxip##JU9g_>~0a2 zw?$MmDVkTUyc@N@q|#~-V@w{8W25HgfN-@Q9{S}+>z)X4exeuFXoLfA-!+)7s)!dm z?~{b35>QnfN1IdynUjM3NX&QO6sKa0N|u73y?3d}`cPtUN*#3bU3j)ub+-)Z<-&b73mMj7UfL3z>tAMb3c*YjlIS7{~P|!=ct9s3@|+(=VpM=EXbFwzQ5qd^(JtBN;F~vl8cg z-GU+G;_$KrfBunL1U5r{Sec{+TSt`Qkq=R5YFP|f+axjUcm*1M+=e0#3-Q~r8r-)r z9Na6DA$G1SC@7AEJmZC^a8MM}rx#<8wlwotIv!=lwNjnhJ$PhyKBRU6lv{Bb`R9_X zN=Ow9U+coB^Y-%|nWdAc&!SX2^$V{`*8>7VW9S8iO~lPn51JY*na($FsB`E7Iy_+< zOtKn7q^^#I0)ZZgIXf2fR1OgLpZ8H(eIyi}x4=T<@32fck&Lm?1kujr(7REEcf8q) z-mBjWW%pK7-F#_YZ(th<$|}U=TZ|YtZXOoRH>3jhPvh3Xdn7GK5Z7IvM5g-1U|RZH zdew3IfGGBuUd%p)cymU~ru0p5uarE4yMNl0h#Gx5^q-+&~ciLA&pZH;1 zKJFw(CBFh>{hwglD+3zsmT9x!YyvLaej6@F`M{iCKH&3e2U@&I<%C{Wgbf7+5NtrH zr%*m=3g3<;VXw)fJ6>S9w2xTZml4}`5meLMK|e1YBG=A-#ja)cXymX6yq(>k(jW!x z?=(`meoM%#j^j-Ebptv>%}{*bb1<2)n~GjYfd%hO@tWLB-WI`J@}@rlGQL@0RB#b+ zGIG#3r4QyuDqzc+E%Z&N8C)&BgR#Q*NZgJ#eDF$`nRe$5kv*tRT~?34re#M^s^%^4 ziF*T_`l1RO_NDM9*9(H%w>cQK@@akFLUH!`8y}o+2bfhg1!rFG=X}_%34UwFgU&}E zNc|E?^N$>)zlJtMGkI17F?PK2}f z)Y-&@Xc|*xjHXU!$cZttII*2aU{pXQ=enyem^Ilzd4x5UTsIOOl9r%*vkdddy`LKE zWPn{cgI~5BqA$+aW6d6Sa6DRpYHMWiQbRaQDBgyjinQ2{P5P)Y%ZQbHB+kw%QDQBv z8t6dL9yA&w!A9(qL)+%Zm|@a_8%#{0sj>ugcgxYUW5Vbb@oE%*7l(_<>Bt?GR$2d&1$GQ;*~J1 z`uIEX_4Hxf+K~cVM@_Hop3Y&uR{QY$H&kbagvPPZRc{L_==W?FVf-5}lJT-Q3 z-FwdSNMp#3IZ*%9?g?lbI-#a}J~e(l2Lwvqf<)*JvZlI?zSKWT@>d-pAS$D|P?UA_JJc|@Pe?fcqpQ1mm>htnHd;%>~6`mR0 zPnLi6BG%>4VTXAkr-t_mebx(Nk)0%VSj97(JENKIuTEHTV<%2*lflbXjVLionW-+$ zCqo`7I8OE(grC1fgUseY!|hxQl=#SbSTGv;vOJlGHM%yFJryyAZ*wlAv+#7yYw9ND zM3T2JhNc4**p(uP`fm;B-QcxoNpx^a=@p23Fp+sO&72iA)4~I9=fkzZc8nZ95@+kw z^Msb{h1?y^%;@9S!TgsiwrusGU)%i9WA`3JHIE;CG`U`o`=K0SSPfZ?a!HWXo$=^vk~T8 z>>}Wqf_qj?ppP1*(OxbbOS96z<*GA>5x9#_xg5@#z%P*1nT@OH7X0v98nmxVKpb~G zns}_B<ky8aGa3h-MZnf42IcM=b7aRxLf(sd&c=^qu@Y7I#lRSV+UuvqxbYrAnf2^bJHn^sPKhx=3-FBSmKiABXmsXOlT;!WPJ-= z8H3PI^z%M(WE)&zyY*DGo1x6kbCbbG_e<-m{GX6RLx!-t=K-yg6oiEz7&62Pz((x@r7oamf>k2miX10`WDW_sCjTKLx{|cvZt?N0`QXufBk2vita8Pa#uSa)~c4)*H4+s_Rxj& zBbkm1ZQYoIy8h4-rH0OtuFT_ulTj>eI$S&T44YmX!Ph!j_FY8+++X6sboRZ5%Gi&1 zsrWt~xo^gMVe=l6&;G#vz9V>UOd_s5{E%2Kkz^MP?1YIE!eL(XGl&&m&Rm(Sz>NNh z!1yLXL_URSac6OU&O=UV`zCC6$fpZNYcg#Ml*pAk))4ga8>C7|(HkvwF!IH9?Bx`Z za5Iph|op^T}v&Xfc zCMJBsyr*{TyW!$G@11Zkp*(JtH)e*^vNFGBa`Xc+K*1-ve08dR5nnzhXw zE1kF4d?OHa=Cs#q7X-uVRa(sKo_jQRIGJOUQABrrwx-{X7{bpOnyM^HcnJsUBOMi(xT)hz5)xAk(*+ozSn#N*aA64r;%ibd4>5(S$4iPy@HWb%=xyw}A?v2P%lx_Niwi4h}!7Avzh2b<{)&(}P!q5ClI zEq}eyBmx2(W03RF7S=p|L4L(pW8@VN6gzwXO`ey7scbPiI98Lpa#4`+suPXgc+<&w z-zZbzgsu3xiQZEE;mn^GJU0Hs&Z11+&O*+Nw1l*IQ+<-%Eg1U0aDy{ZGPG zuSA$8qR0gF2a>3Z>8Ko`L$cgoD+uNRxdFLjy0bI<5i*zc`*u=J}&0jh1AkT zoy*9FxHBlFaRNsuou^`@2RU*JuJh&(wUU~Sul3=&YdIf`^x*qk3y?M4N^FPDQ=#^= zoHd03#CxVVopfLdSX*o$Z*O(Nronsop?U{(Tptcj(99~A3$(vH8*CW|^Pb|T6jwtBXe4)-GTVU(i z7`iIU71VA&$LMEjjP#CNn6ojJ>>W8wv!cpjqe~E!%+p8h$)n+N?>Rzti89Ba2?kx4 zfws2;%sD%QR_bhoB?bk=N2~xZiFf147rFT4_z!TIq=hXjRq%q*OK`Y;44#Ed!kVe} zAUs~A{yhl?{{uQO%l`r9+a?l=&zTUtUxBl_A&Ga$X)WX*J`auNIq;!Sj|pC~6qc3t zKo55zaa9*%J0I8L&ACw!e(e%%d%2KPIi$lJ(WvD4ULU4P9enkaY|+ev+iQ6=^r9l!T^tSb(^rDt z!Ekc*!fTrTX&eMPg>fY zgdJn1vzp9Iy$d96dMvIyrcbqkOj&(3V`kKTCu|IFL9aoM{Yhc8o&H zX5e%D3L(J<9(B0)l z&fCp|A60vx>xwOZuGHtHtP_I=f@X|m`V?BSZXK=&IEe)=$-LWUsg%=yh`7Czr@=L0 ztjZEy%8k@wnui8xOr#Rp?TzBiaVjP=)eoZZH3=Bun}ba&U&HpywXl190=Q>3aa!N! zkehFo+t|&x0rnpkplZr`R3G<_)CsMJ%N4cYS;%1)6fPu9j+29aMw zQ?U8dUUE33fu@L#XZohHywq_W_3Ithv7dDem`*)P+!vRS#Q7Ymx5I=yd&?nqx0K04 z!%SLvD+;UamqVCI0LbyzYC#WkXx)B7*wt0Uf1Dtkepf~_t-Wxpm^du$8AD8Gy5i|o z1ZK^c%=CO67l%y38|gBdV&L2CC-_Fr;J)v^JloeIFG-Wy&|~ z@h^l-jioTg#*oQ26l3NstwA-TLE5N6QAjkNXML>-JHHCyMZI?15GcsHJuilLN2}0R zs{(hB=@5NWgngJzz#&eNJ)F=%v>QTcp1cRh{JcRvi5$cIW9w*&Q8#%I6wS#Hk0gfE z_mXdmJjjE>GCE804e8q%OC)Ao1n3oHw@jWz%a3ozY(rNZP&^JTZOcI{#|EaguSVT1 zNhoq}Hg3o~OH@sBY0HCVbeyS%k9%uj1e3`dTpC3!B`ff9oG#{hTt?fit6*R0df=U% z#DL0I8|90I9KDYRvD&Q&U6PB@`ptK=h!A23FrHkr-38(DLqx+U2r8QmnOC~8VANiZbfE#ScVjLuptQWcx1tKCt)2)c^p4`) zn~L}-a01((--mWqN$BoYfLF3B!SRF+K5CU^>ZH4PVRyw@)+&r_+n|Nt?=FSJ)*YOR zSOHvjY6)Zei_7j248)@RZmPd76P$JQnZc`@$vN3V;^A_XPBV_7!pH0B`A?ph(i2NX zt2?Q4t_*6Q9j7f3nZ|GU58G%tFZGwUE(O0D+Q zL3?5*PM_RBikX@CNv)2Qs~X}cpLaa*gi3U>Rb`_4h3Mc{1yHgVLW>(Lr%QbvYIMpo zO&Xpw!^at`W6W8J&BYu)Uk|ulV$S)oM;K1Y-h_}aEv)F_`{>IZZOkuU;f);HO%K1g zjO#}xlhc=m;X`UPjo$PT_i7kXt(lSN)*;4rDDP$a;-8_wua96|R*43?6nXD0?ok)P zI(WZU3@0pBfr*w2$+qMCjc@yQu)Zt<3)8D0G5aTuH4(sB7uDD<17Qp>RAb)G_5x?V zZAUNLO(#7Lp^FaAf&?pnHYg{FiIZ_8)z?puMLUP+@z`Q&YLE;2)u)nJ&Q`F)7O=ck z0EYr4L2OzmXa9snU{I@#k_XS>)j?${Cw_xtJ@_2x!B4=MDut77l#w0Z35rG>wc+SK zBjtl>&=_EifvNc*zimAouP6y`w~MmzoQ1qg;xp)l$}s3Fx8dEGd6zD)T}yXA$)*n* zRG8rl{J3PV7;3CkA#+kEP_}9ebSpK|r1ASeW%&r&H}(j87@^L_NVRcNyiLd%fi2Wr za3U5>*Cf8j50EPc6YggmaLDUd=_k)xNGk z+IyhcF^o11isRzBd2Fk5YW?HalF-mgh>xKb2atgw0vU^d7i4p$JV#;Y3;0-Zlz7>l;FK7w0D;f@IX54T1cw8L zaPp-YC@oNB1?Ai*>#_*nPIyU+r`qlYmj@Jy|Jhbv8885;pF9 zL}u7}L11Pnw7bf&O`>ObmpMz}M}P)vy7U=6UsA?Nv40C=+NF3UHT$T|(uFLq&5l3I zR1lj37Ubf&CuFOQJZ9g?<8&sK)7WPVU|g^f^zX~3#Ss(WW$8v7$yfldM_ej`(teotD3=#kWoVCH+Gyoj^nUP z_X_$&jbszLFYsgXg>d1eB0Zn~omf0+<6YwWf8(#5rO$*^+3KXpVAx`Svwzv+XP?L5 zT(AsBa9)xd?xS(q5I+`a(?{#CyrBvEH$%4LbX2lafy&aeACe~4H#4anKh|VheGu$Dt2pC3S774WLUi9PgLC>Ec=pz^P_%Rlc07%z zWo>onJ(KSr+`CI;?E#_=RqJ}0LQ7eijoDKu8GME3qq zTC(UIS^0B3KQ_M|b}pU_S5!;!&;$$I;ipggvOdxP+c=($P6jQ^+>4>Rim_4H2G6d} z#T#?45W$_%aPv?ZE#7C#xRvhV?45Z6Ru1lgvetvpH`ohHZ;pasQ8{qly$J*P`>!=~ zSAk#R58kvfX5f5%7UQdT4fDrv>kXI*?2@n%jG0CNOnoWM$ca~yyvksT(GxJQ{uC6K zJwU@3pD^OdDqK)C4G$&^(U52f>QWv}#6-{Fo+CX_J~7EgU7!<&rN?nJsXO*gnnrY< zCDJsh>%?OC3di~B3lKYY4d+>j;`<11SR5)wVv>jH`jXG2W3k4$Lw6sO&`~4d_M+P` z`gk~5FwPmIR>}gioRAUG@znH|CS2MyokUlDAZM~Ap(j$7oF4cMLVy}Iu!hdCZ zPEp1CxBAIolLRz8HiXd=rZGQeJmTbizk{ucqDX{0IFD+A@O92yjJYDr6T1!6zO;m= z_~sg>&KSj3b22gUo)=yo+yMuaWk|zZSw>cGm$h%pYQEniO#>wN@V+|;L(1)X(z`qz zCG&>Bc&ZjV#ykiOYhy@*QU@=E*-0{!KLFe5j|r(o#C(P^Q&yz}yeSu;h4ZR@|9L6K zMMNBzNTlKVx6V{9dmP3!Xff*Aq0G12c^Dt_nLZDBg5|R^>B_^GK(?_9^tP^|K0fK7 zzaM!QY{cP=-2=R_eItUSBJq9K4HF|)pm@9*>I?aTno=dTz zf88~~tc^s4XcDu8GWvQ)9E@4Lj2Wfm!Hz8S!tI}qgU+|JRQ{|iS?(qewCFZ=e|OpiH?JtH!4u=_YnHCVuL_3^=w@H@ci zQ)S%Woq(%h`>|?g1Wk&-(cYKH6-YsL#JuoI2 zO-`)SvW-OWayJ|;oXLzF$ikW#nMftxlTSlYtfjCLeR#~2_8;5A9M}@fD&$w8LCt<# zji-qNrvO$t4)VC}2k52YP^|2gV0=Xju>boN`uda%nr1$v3$~n~n-${mTX!%V{Lp73 zYZ?i|lDXu@peS;itI1*oUp_}?0(-7X6^CwZCC}5|(-YeD#Lwe2f@ljWEbf5|_x9kt zHw3P&dd@K&_X~tvHsSXt-F*7VQQm?&15hcIgGD36m~k_MIiFvq;{x4YTsh_hS+csD z&kty(b9EnrKu`qUD187OHIX?+`MaR0356BA~eL3hF~PE^v;*H5Lz;LB8sG&3f3Pm&csD)a$$Cg5PW{32%dUYXcd%0KHiwZ3Vv&c+aDIAex?NT zBIhynJjTIW0SWl(vIJ!E#xtKvI8Yo=LAC1oabk8i$bOm3^WVK2^c_~Pl5w_hF8($a zRmw2ixWkw}Hx;5BmgDw|^XS*ZO04IrQzSS=20HT<$?dKMj7IDP=A+|PQ0OnhC639Q z_hZJfhFNbpY4Le*aFac?`Xa^*>{*B3n+5Cd{&>sllY5FTQxiZ#Ns$#inLrGU?m+pA z(d_e28l>-(1D%_g0=}is=vnVi#IgDUjTp?O`!3{R^*U+LFFrz_A3YCS7)>;|G!awG zD`KK)Rrj{xr@7^vWU?+OeGTr(Ceb>NsR=ALKDG-AZ*6I7TO>#dl>_xxA`7)AL+s=UOjSw_oRz4gcfabA z;4=XrZsH5FUzdQ~Nk!D=)sWcodLk6#4a)DzVAruT#Mp2*%?TS1F{yi@V%ZSqmPQbr zYGVwmhvGoLU6Vw0JtntfVyS&{BnW&eqQ8c^NY*8B62RXVEB#c<5jy{kchy(}o7YwI z0^cj}K9nGJ|0xbC8>;B=-4goD*p^q2#*%P@*{B_z!ZNvDrJM=Vr#ku7aiwKKZZ z$eNX-P&+9Nr(AHu&mrsSnKMePMUoI^Ho1Z3jl*=goC6+k?jeo~Mqrr8R`RVRj5l3s zC(f83PbCf&V)u7>9M?4g4s;cfg^u0mKgShE+3kW~zLJ&>pRipVWJa*M|XGtGowRoIFK8mia&vI*^HN-$;eO zDGt1GhPEClyyS~8wbhcgo~R_*v%GLl>;dbr(TQN%c!ubeHruRzKY?A?eiqgQOd@;j zMd62;FZ-rF8m6VzAy@e+-0`nLl{dpg#_$l1WRKCwE*=v6Z!RgTB}B{=4qBCd)S!=|q_beX+8?SH%)AMO~> zTT`0}S7YX)L-LXu1}n*fJ?p|0Texs@2fu z$)7+??F*5!qhP8S4$BLcK~Mn)L%dbNw4Kib+uA{epD2*gC$^D?=dJO@3R{?-osBzl z_kmxa8B|11htt)MX+&EL?+}>}AFXrg$;E1Dy(f%(@=c;vd`Dt%&Q@9%`HC1`Z^4p6 zdDyhJ40g6L%)G5qIEN$3s&6RZrKn2b#yK^((|!!y>U9c#9?<~LPh8%J=cAxzkt4Kj z{z~_^_F~+r1n>%M1m89DAZB3-goyuw@~?5sxvZDq+wcyR2QpySO9xoD{|dRZ`7Ubt zH_^C@(NtRI0A?RnVD}Ia7ISM+`@%=MIA#cPKWf9~m~NQ*GYcxG39^ft#Bt&482o5> zj<;SY1?-YoIx5zm$-5iER`yiV)QFkLOU)#O=9#?1+ygKv@Eyi|;lOT#jo4^ofM*;U zh^M22H^L=>F9>OoapOsL1l&(y3Nk2*YwQ7g?Hoe zc~mQ|N_b9MCuKwb_;yGnOfLHjbYNp~4L0 zRVicVxJN`DzM|wb8PX|ykT^GZp~#dhC=1($&n6|1GTs^91vM%5dbJ;{(5SXf$~OY* zD`}{BU?SS;zsDzEKl7Zg{-iw{RB+nIW@_`*1s5NhNrHzw07JLpMpY%;=nQ1b=Xrci zYC5slbOxh!J2=ZsJBaH$8~7rUiLO(#aqSypwm(@NC%IE--}R0sZE&Nm{{(;iU~5KO zG>(vyw1v+2p+&MT*b}?pWLkJr9aVGpV`Py$i1ldW)Y`ppsdX}ZzY&h3wA69K#y!N~ z{Z?o^6orjXM^eS17rgzu!_n$O9Uf{l6sy+s60SFEA$>YmZtoIYzb zS%3?quamW1iCAO8gLpm%a3@zB$3Gf{aqlKlwMG+I!W+wJA1)<6U*_Odk^M9ZGw?*q zdKwon4I_=0fLLsLxI-ueC@^-`XF8+@96lT`<$Lp$K!L`hhb(2MDZ;&Ipc zSURIphtYf#g?IWZ@X7dH`08g8BN%7}qm~{6Z)sU}N%(S@n|lPgx#!@aaS$B$C(OHZ zxxBzJain*+nM_lf$lCQ>#QBpX*(Vwz%+ug#sBPc8{R=<3?BhJv(!fPKWZ{m~ z3-p`33SIj)l4QMT(3)??@{~W(BYHmc*a9tPFP{Uqy;A_IPrRZ_%ny-(9bwdA(=AZ< z9<*_s;*Nq#8aTV-e5lLCT*%$No8G>Wgo#6O7~-D|8?Tx`ltw++Ic9^@D`zG-T8h~= z)PUV@^T?-7;{4}Oj+{@kMsm6ou7w6-yfP`I z@l`3*_H{FnfA2xehKg`nfeW;z?*-=P4r*XH0}nesCnDnEU~`c{scKOwmt2Y?2TJ(- zgg{7rC5>YZdAyj!Q?R4;E(uF;1J5gYaP7hh{(9jZ6>)0h%+-@%*1wP;W|2uyHZ6c; zxE=-l_2#I$Py|c+H<3qAHpBj?LMY7gfU!ZF(O*s#e6+oJ#}rEGMUkE06E^}H6m)Rg zNd@SeZcE?t_jwnpBw5S+^uQF&c6#3?lh{tmM6;d2_;}tkYPT>6nQvd|%G1kug|lKX zw51!3mb@S(D+}m#bVcWnzev=Aha{oJ6-uAUvk%i^(IG>aM6S)laq`ZbOGv^x z{@UWp{&aF~cRUmYWWiJy5$0|85!ipZ2y|U{kx1vmMDa!fp`DwtR{AFBs>RXlQZ4*p zQv{jn&*)(fOT4%83ARp*1kR*7vf*4heYN8&@e*sK(mH&<>*rj&Zt)He>+C`gjm1#% zx`(I_h(Kx|1Kg%7I3>B1Rw+G%2W?kr$#+3^Q$`2bu9F2}<5O|`_afw;FlV$koyJY< zH6j}Mn6}-IMq8WFxFTvRF?rqo7Y}dQva`eT69oi*m;Z@}XZ|M-Z;HL`<0R8{gTK+U8R4Gf@~VJqt?wqToa} z3(L0cg^0~6P~9ScH3t;&^KJ{eRC*##9UF#`%Zs5v#DPeD)j}^mpW$`UD_&!Q3<&ah z63UJVut$7=NUc)F`C@VqPaRR>!wd33TMB;j7k4bFQ~bZ+FPi_q@E4zl_D~1j0PN2{ z4MyU-$!5Dl*fmF-nJ_ONWz$5kub~DV`j27T*ttBR+IWt>&{5Kcqaott17dP+IV^Ar zv-#XHg-P7}o(g|QRNt5Zw(_=|&RIIN-N_vNr_RJT8mmG5mKI0w*JHFUG6a_&mry=k z2nM&vk)DR_#Kmq67S7!aRWn-P?7jqCGRubZ$SfC6$!cP2k2Z}}HOD9G58=hMe4Nb4 zQXRc*uo(dmLfKA_B@C4uYQ2Q_!1m%W7s?ZX(mwMxjPZx2#EpTM_ z=t6o%I?S|Er1|U8;Ni1;>Uw)GR+4$}N?aJ5e|eI7Z5!zQSY!Hf_X*xN{@xbHJrU-9 z^+AL9SXk!eiBoKa@MqT&+P>%r4H>CVw_jsXbEz~hD*P~B4{*UlvyT9N-inX>vdL$? zW_W7h&VG5k1Tw-05oek~y~9NqSUHSuBt#hVy4{R#tru+O>awlFTCGj?am_$P;u|FQcw$RHH{=v1pp)hZ zRM;m(W~`dbE*EKrYf+bAtO>*Mo0^Nf+x*zYq8fbJ*bZ03jfqL;TD%#0iNgpjWq)ln zfZ~>JTCt=QA8MTf_V#3^URezH%gC~?-nUUpk8y;JxlCGbe8(W&vC#Lk9X@}r1R-rJ zxH`EOM|^xpPiCp2VkrH$qeg-#BAB4N&Z(u@~DMsWL!kUE={Ogrs z8!mjH#j#c}Teg!=qVpuDOVW5z6%nvr=On~$E(PTYp=9OltI%63MlG_6L8o;y>WsYr z*I$3f8(S}9>z21!$ST)t=WY8}hs1;#9mfR#lQa4xG#4JiQ=8e)q-HYNROgdDy?pf8Up!fXTo6`mfYT zdt1Z*wzGe;uK)L!DU!c?Ydf%d!S8I!-^JZ&#P8F;b9(N&>~f2LV5j^O z`=2)QCpLHLKd|Qi#Qx_!{}Vf4{SU0gKe7LL2j%`4=an}98~a~2;_f6WYV==?w1g)8 zyZOJexc_y2_B%hCf864W_|M$E7-);XF!(ccj literal 0 HcmV?d00001 diff --git a/evaluate_models.sh b/evaluate_models.sh index 830752f..5f64408 100755 --- a/evaluate_models.sh +++ b/evaluate_models.sh @@ -13,4 +13,11 @@ python -m src.eval_main # idm python -m src.eval_main --method=idm -python -m src.eval_main --method=ogail --policy_file='checkpoints/gail-options-setobs2-15-02-2022.pt' --env='NormalizedOptionsEvalEnv' +# options GAIL +python -m src.eval_main --method=ogail --policy_file='checkpoints/gail-options-setobs2-Feb15_18-49-05.pt' --env='NormalizedOptionsEvalEnv' --env_kwargs='{stop_on_collision:True}' + +# options GAIL-PPO +python -m src.eval_main --method=ogail-ppo --policy_file='checkpoints/gail-ppo-options-setobs2-Feb15_22-05-38.pt' --env='NormalizedOptionsEvalEnv' --env_kwargs='{stop_on_collision:True}' + +# behavior cloning +python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' diff --git a/src/evaluation/evaluation.py b/src/evaluation/evaluation.py index 3682759..20a6db9 100644 --- a/src/evaluation/evaluation.py +++ b/src/evaluation/evaluation.py @@ -122,7 +122,7 @@ class IntersimpleEvaluation: done = local_vars['done'] _agent = info['agent'] env = local_vars['env'].envs[venv_i] - assert isinstance(env, Intersimple) + # assert isinstance(env, Intersimple) self.eval_policy_step(info, done, _agent) diff --git a/src/gail2/envs.py b/src/gail2/envs.py index b9cef2b..6f6bdd4 100644 --- a/src/gail2/envs.py +++ b/src/gail2/envs.py @@ -25,11 +25,18 @@ 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)]) +def NormalizedContinuousEvalEnv(**kwargs): + return Setobs( + TransformObservation(IntersimpleLidarFlatIncrementingAgent( + n_rays=5, + **kwargs, + ), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) + ) + class OptionsEnv(Wrapper): def __init__(self, env, options): @@ -65,7 +72,7 @@ class OptionsEnv(Wrapper): o, r, d, i = super().step(u) actions[k] = u rewards[k] = r - env_done[k+1] = d + env_done[k] = d infos.append(i) observations[k+1] = o @@ -84,7 +91,7 @@ class OptionsEnv(Wrapper): ll_obs, ll_actions, ll_rewards, ll_env_done, ll_plan_done, ll_infos, ll_steps = self.execute_plan(self.last_obs, self.options[a], render_mode) hl_obs = ll_obs[ll_steps] hl_reward = (ll_rewards * ~ll_plan_done).sum().item() - hl_done = ll_env_done[ll_steps].item() + hl_done = ll_env_done[ll_steps-1].item() hl_infos = { 'll': { 'observations': ll_obs, From 5bd8b42d9fa8768ce0ca03ed4638dfe7418095f3 Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Thu, 17 Feb 2022 22:41:55 +0100 Subject: [PATCH 05/10] Merge updated files --- scratch/etienne/trpo/.gitignore | 232 ------------ scratch/etienne/trpo/core/discriminator.py | 74 ---- scratch/etienne/trpo/core/gail.py | 125 ------- scratch/etienne/trpo/core/optimization.py | 39 -- scratch/etienne/trpo/core/policy.py | 96 ----- scratch/etienne/trpo/core/ppo.py | 72 ---- scratch/etienne/trpo/core/reparam_module.py | 162 -------- scratch/etienne/trpo/core/sampling.py | 73 ---- .../etienne/trpo/core/test_optimization.py | 23 -- scratch/etienne/trpo/core/trpo.py | 79 ---- scratch/etienne/trpo/core/value.py | 48 --- scratch/etienne/trpo/core/value_estimation.py | 40 -- .../bc-intersimple-setobs2.py | 2 +- .../gail-intersimple-minobs.py | 2 +- .../gail-intersimple-minobs2.py | 2 +- .../gail-intersimple-normobs.py | 2 +- .../gail-intersimple-setobs.py | 2 +- .../gail-intersimple-setobs2-recurrent.py | 2 +- .../gail-intersimple-setobs2.py | 2 +- .../{ => experiments}/gail-intersimple.py | 2 +- .../{ => experiments}/gail-options-minobs.py | 2 +- .../{ => experiments}/gail-options-setobs.py | 2 +- .../{ => experiments}/gail-options-setobs2.py | 10 +- .../trpo/{ => experiments}/gail-pendulum.py | 0 .../gail-ppo-intersimple-minobs.py | 2 +- .../gail-ppo-intersimple-normobs.py | 2 +- .../gail-ppo-intersimple-setobs2.py | 2 +- .../{ => experiments}/gail-ppo-intersimple.py | 2 +- .../gail-ppo-options-minobs.py | 2 +- .../gail-ppo-options-setobs.py | 2 +- .../gail-ppo-options-setobs2.py | 8 +- .../intersimple-expert-action-profiles.ipynb | 0 .../intersimple-expert-rollout-minobs.py | 2 +- .../intersimple-expert-rollout-minobs2.py | 2 +- .../intersimple-expert-rollout-normobs.py | 2 +- .../intersimple-expert-rollout-setobs.py | 2 +- .../intersimple-expert-rollout-setobs2.py | 2 +- .../intersimple-expert-rollout.py | 2 +- .../ppo-intersimple-minobs.py | 2 +- .../ppo-intersimple-minobs2.py | 2 +- .../ppo-intersimple-normobs.py | 0 .../trpo/{ => experiments}/ppo-intersimple.py | 0 .../{ => experiments}/ppo-options-minobs.py | 4 +- .../trpo/{ => experiments}/ppo-pendulum.py | 0 .../etienne/trpo/{ => experiments}/readme.md | 0 .../etienne/trpo/experiments/requirements.txt | 3 + .../trpo/experiments/sgail-options-setobs2.py | 108 ++++++ .../experiments/sgail-ppo-options-setobs2.py | 107 ++++++ .../trpo-intersimple-minobs.py | 4 +- .../trpo-intersimple-minobs2.py | 4 +- .../trpo-intersimple-normobs.py | 0 .../trpo-intersimple-setobs.py | 4 +- .../trpo-intersimple-setobs2.py | 4 +- .../{ => experiments}/trpo-intersimple.py | 0 .../{ => experiments}/trpo-options-minobs.py | 4 +- .../trpo-pendulum-rollout.py | 0 .../trpo/{ => experiments}/trpo-pendulum.py | 0 .../trpo/{ => experiments}/trpo-walker.py | 0 .../etienne/trpo/experiments/vec-env.ipynb | 346 ++++++++++++++++++ .../wgail-intersimple-minobs.py | 2 +- .../wgail-intersimple-minobs2.py | 2 +- .../wgail-intersimple-setobs2.py | 2 +- .../{ => experiments}/wgail-intersimple.py | 2 +- .../{ => experiments}/wgail-options-setobs.py | 2 +- .../wgail-options-setobs2.py | 4 +- .../trpo/{ => experiments}/wgail-pendulum.py | 0 .../wgail-ppo-intersimple-minobs.py | 2 +- .../wgail-ppo-intersimple-setobs2.py | 2 +- .../wgail-ppo-intersimple.py | 0 .../wgail-ppo-options-setobs.py | 2 +- .../wgail-ppo-options-setobs2.py | 6 +- .../{ => experiments}/wgail-ppo-pendulum.py | 0 .../trpo/sb3/sb3-ppo-intersimple-rollout.py | 2 +- src/core/reparam_module.py | 3 + src/options/envs.py | 111 ++++++ .../etienne/trpo => src}/options/options.py | 17 +- .../trpo => src}/options/test_options.py | 0 src/safe_options/collisions.py | 185 ++++++++++ src/safe_options/options.py | 305 +++++++++++++++ src/safe_options/policy.py | 35 ++ src/safe_options/policy_gradient.py | 107 ++++++ src/safe_options/test_options.py | 54 +++ .../etienne/trpo => src/util}/wrappers.py | 0 83 files changed, 1439 insertions(+), 1123 deletions(-) delete mode 100644 scratch/etienne/trpo/.gitignore delete mode 100644 scratch/etienne/trpo/core/discriminator.py delete mode 100644 scratch/etienne/trpo/core/gail.py delete mode 100644 scratch/etienne/trpo/core/optimization.py delete mode 100644 scratch/etienne/trpo/core/policy.py delete mode 100644 scratch/etienne/trpo/core/ppo.py delete mode 100644 scratch/etienne/trpo/core/reparam_module.py delete mode 100644 scratch/etienne/trpo/core/sampling.py delete mode 100644 scratch/etienne/trpo/core/test_optimization.py delete mode 100644 scratch/etienne/trpo/core/trpo.py delete mode 100644 scratch/etienne/trpo/core/value.py delete mode 100644 scratch/etienne/trpo/core/value_estimation.py rename scratch/etienne/trpo/{ => experiments}/bc-intersimple-setobs2.py (96%) rename scratch/etienne/trpo/{ => experiments}/gail-intersimple-minobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-intersimple-minobs2.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-intersimple-normobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-intersimple-setobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-intersimple-setobs2-recurrent.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-intersimple-setobs2.py (98%) rename scratch/etienne/trpo/{ => experiments}/gail-intersimple.py (96%) rename scratch/etienne/trpo/{ => experiments}/gail-options-minobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-options-setobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-options-setobs2.py (89%) rename scratch/etienne/trpo/{ => experiments}/gail-pendulum.py (100%) rename scratch/etienne/trpo/{ => experiments}/gail-ppo-intersimple-minobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-ppo-intersimple-normobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-ppo-intersimple-setobs2.py (98%) rename scratch/etienne/trpo/{ => experiments}/gail-ppo-intersimple.py (96%) rename scratch/etienne/trpo/{ => experiments}/gail-ppo-options-minobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-ppo-options-setobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/gail-ppo-options-setobs2.py (89%) rename scratch/etienne/trpo/{ => experiments}/intersimple-expert-action-profiles.ipynb (100%) rename scratch/etienne/trpo/{ => experiments}/intersimple-expert-rollout-minobs.py (96%) rename scratch/etienne/trpo/{ => experiments}/intersimple-expert-rollout-minobs2.py (96%) rename scratch/etienne/trpo/{ => experiments}/intersimple-expert-rollout-normobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/intersimple-expert-rollout-setobs.py (96%) rename scratch/etienne/trpo/{ => experiments}/intersimple-expert-rollout-setobs2.py (96%) rename scratch/etienne/trpo/{ => experiments}/intersimple-expert-rollout.py (94%) rename scratch/etienne/trpo/{ => experiments}/ppo-intersimple-minobs.py (98%) rename scratch/etienne/trpo/{ => experiments}/ppo-intersimple-minobs2.py (98%) rename scratch/etienne/trpo/{ => experiments}/ppo-intersimple-normobs.py (100%) rename scratch/etienne/trpo/{ => experiments}/ppo-intersimple.py (100%) rename scratch/etienne/trpo/{ => experiments}/ppo-options-minobs.py (95%) rename scratch/etienne/trpo/{ => experiments}/ppo-pendulum.py (100%) rename scratch/etienne/trpo/{ => experiments}/readme.md (100%) create mode 100644 scratch/etienne/trpo/experiments/requirements.txt create mode 100644 scratch/etienne/trpo/experiments/sgail-options-setobs2.py create mode 100644 scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py rename scratch/etienne/trpo/{ => experiments}/trpo-intersimple-minobs.py (96%) rename scratch/etienne/trpo/{ => experiments}/trpo-intersimple-minobs2.py (96%) rename scratch/etienne/trpo/{ => experiments}/trpo-intersimple-normobs.py (100%) rename scratch/etienne/trpo/{ => experiments}/trpo-intersimple-setobs.py (96%) rename scratch/etienne/trpo/{ => experiments}/trpo-intersimple-setobs2.py (96%) rename scratch/etienne/trpo/{ => experiments}/trpo-intersimple.py (100%) rename scratch/etienne/trpo/{ => experiments}/trpo-options-minobs.py (95%) rename scratch/etienne/trpo/{ => experiments}/trpo-pendulum-rollout.py (100%) rename scratch/etienne/trpo/{ => experiments}/trpo-pendulum.py (100%) rename scratch/etienne/trpo/{ => experiments}/trpo-walker.py (100%) create mode 100644 scratch/etienne/trpo/experiments/vec-env.ipynb rename scratch/etienne/trpo/{ => experiments}/wgail-intersimple-minobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/wgail-intersimple-minobs2.py (97%) rename scratch/etienne/trpo/{ => experiments}/wgail-intersimple-setobs2.py (97%) rename scratch/etienne/trpo/{ => experiments}/wgail-intersimple.py (96%) rename scratch/etienne/trpo/{ => experiments}/wgail-options-setobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/wgail-options-setobs2.py (96%) rename scratch/etienne/trpo/{ => experiments}/wgail-pendulum.py (100%) rename scratch/etienne/trpo/{ => experiments}/wgail-ppo-intersimple-minobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/wgail-ppo-intersimple-setobs2.py (97%) rename scratch/etienne/trpo/{ => experiments}/wgail-ppo-intersimple.py (100%) rename scratch/etienne/trpo/{ => experiments}/wgail-ppo-options-setobs.py (97%) rename scratch/etienne/trpo/{ => experiments}/wgail-ppo-options-setobs2.py (96%) rename scratch/etienne/trpo/{ => experiments}/wgail-ppo-pendulum.py (100%) create mode 100644 src/options/envs.py rename {scratch/etienne/trpo => src}/options/options.py (96%) rename {scratch/etienne/trpo => src}/options/test_options.py (100%) create mode 100644 src/safe_options/collisions.py create mode 100644 src/safe_options/options.py create mode 100644 src/safe_options/policy.py create mode 100644 src/safe_options/policy_gradient.py create mode 100644 src/safe_options/test_options.py rename {scratch/etienne/trpo => src/util}/wrappers.py (100%) diff --git a/scratch/etienne/trpo/.gitignore b/scratch/etienne/trpo/.gitignore deleted file mode 100644 index 7491629..0000000 --- a/scratch/etienne/trpo/.gitignore +++ /dev/null @@ -1,232 +0,0 @@ -PyTorch-Reparam-Module -cg.ipynb -vec-env.ipynb -*.zip -*.pt -*.mp4 -*.pkl -runs/ - -# Created by https://www.toptal.com/developers/gitignore/api/linux,macos,python,visualstudiocode -# Edit at https://www.toptal.com/developers/gitignore?templates=linux,macos,python,visualstudiocode - -### Linux ### -*~ - -# temporary files which can be created if a process still has a handle open of a deleted file -.fuse_hidden* - -# KDE directory preferences -.directory - -# Linux trash folder which might appear on any partition or disk -.Trash-* - -# .nfs files are created when an open file is removed but is still being accessed -.nfs* - -### macOS ### -# General -.DS_Store -.AppleDouble -.LSOverride - -# Icon must end with two \r -Icon - - -# Thumbnails -._* - -# Files that might appear in the root of a volume -.DocumentRevisions-V100 -.fseventsd -.Spotlight-V100 -.TemporaryItems -.Trashes -.VolumeIcon.icns -.com.apple.timemachine.donotpresent - -# Directories potentially created on remote AFP share -.AppleDB -.AppleDesktop -Network Trash Folder -Temporary Items -.apdisk - -### Python ### -# Byte-compiled / optimized / DLL files -__pycache__/ -*.py[cod] -*$py.class - -# C extensions -*.so - -# Distribution / packaging -.Python -build/ -develop-eggs/ -dist/ -downloads/ -eggs/ -.eggs/ -lib/ -lib64/ -parts/ -sdist/ -var/ -wheels/ -share/python-wheels/ -*.egg-info/ -.installed.cfg -*.egg -MANIFEST - -# PyInstaller -# Usually these files are written by a python script from a template -# before PyInstaller builds the exe, so as to inject date/other infos into it. -*.manifest -*.spec - -# Installer logs -pip-log.txt -pip-delete-this-directory.txt - -# Unit test / coverage reports -htmlcov/ -.tox/ -.nox/ -.coverage -.coverage.* -.cache -nosetests.xml -coverage.xml -*.cover -*.py,cover -.hypothesis/ -.pytest_cache/ -cover/ - -# Translations -*.mo -*.pot - -# Django stuff: -*.log -local_settings.py -db.sqlite3 -db.sqlite3-journal - -# Flask stuff: -instance/ -.webassets-cache - -# Scrapy stuff: -.scrapy - -# Sphinx documentation -docs/_build/ - -# PyBuilder -.pybuilder/ -target/ - -# Jupyter Notebook -.ipynb_checkpoints - -# IPython -profile_default/ -ipython_config.py - -# pyenv -# For a library or package, you might want to ignore these files since the code is -# intended to run in multiple environments; otherwise, check them in: -# .python-version - -# pipenv -# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control. -# However, in case of collaboration, if having platform-specific dependencies or dependencies -# having no cross-platform support, pipenv may install dependencies that don't work, or not -# install all needed dependencies. -#Pipfile.lock - -# poetry -# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control. -# This is especially recommended for binary packages to ensure reproducibility, and is more -# commonly ignored for libraries. -# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control -#poetry.lock - -# PEP 582; used by e.g. github.com/David-OConnor/pyflow -__pypackages__/ - -# Celery stuff -celerybeat-schedule -celerybeat.pid - -# SageMath parsed files -*.sage.py - -# Environments -.env -.venv -env/ -venv/ -ENV/ -env.bak/ -venv.bak/ - -# Spyder project settings -.spyderproject -.spyproject - -# Rope project settings -.ropeproject - -# mkdocs documentation -/site - -# mypy -.mypy_cache/ -.dmypy.json -dmypy.json - -# Pyre type checker -.pyre/ - -# pytype static type analyzer -.pytype/ - -# Cython debug symbols -cython_debug/ - -# PyCharm -# JetBrains specific template is maintainted in a separate JetBrains.gitignore that can -# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore -# and can be added to the global gitignore or merged into this file. For a more nuclear -# option (not recommended) you can uncomment the following to ignore the entire idea folder. -#.idea/ - -### VisualStudioCode ### -.vscode/* -!.vscode/settings.json -!.vscode/tasks.json -!.vscode/launch.json -!.vscode/extensions.json -!.vscode/*.code-snippets - -# Local History for Visual Studio Code -.history/ - -# Built Visual Studio Code Extensions -*.vsix - -### VisualStudioCode Patch ### -# Ignore all local history of files -.history -.ionide - -# Support for Project snippet scope - -# End of https://www.toptal.com/developers/gitignore/api/linux,macos,python,visualstudiocode \ No newline at end of file diff --git a/scratch/etienne/trpo/core/discriminator.py b/scratch/etienne/trpo/core/discriminator.py deleted file mode 100644 index 8074c32..0000000 --- a/scratch/etienne/trpo/core/discriminator.py +++ /dev/null @@ -1,74 +0,0 @@ -import torch -import torch.nn as nn - -class Discriminator(nn.Module): - - def __init__(self): - super().__init__() - self.nn = nn.Sequential( - nn.LazyLinear(50), - nn.Tanh(), - nn.LazyLinear(50), - nn.Tanh(), - nn.LazyLinear(1), - ) - - def forward(self, states, actions): - return self.nn(torch.cat((states, actions), dim=-1)).squeeze(-1) - -class DeepsetDiscriminator(nn.Module): - - def __init__(self): - super().__init__() - self.elem = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - ) - self.glob = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(1), - ) - - def forward(self, states, actions): - actions = actions.unsqueeze(-2) - actions = actions.expand(*actions.shape[:-2], states.shape[-2], actions.shape[-1]) - sa = torch.cat((states, actions), dim=-1) - return self.glob(self.elem(sa).sum(-2)).squeeze(-1) - -class RecurrentDiscriminator(nn.Module): - - def __init__(self): - super().__init__() - self.state_dim = 10 - self.state = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(self.state_dim), - ) - self.glob = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(1), - ) - - def forward(self, states, actions): - actions = actions.unsqueeze(-2) - batch_size = actions.shape[:-2] - set_size = states.shape[-2] - action_dim = actions.shape[-1] - actions = actions.expand(*batch_size, set_size, action_dim) - sa = torch.cat((states, actions), dim=-1) - - state = torch.zeros((*batch_size, self.state_dim)) - for i in range(set_size): - state = state + self.state(torch.cat((state, sa[..., i, :]), dim=-1)) - - return self.glob(state).squeeze(-1) diff --git a/scratch/etienne/trpo/core/gail.py b/scratch/etienne/trpo/core/gail.py deleted file mode 100644 index d630bd4..0000000 --- a/scratch/etienne/trpo/core/gail.py +++ /dev/null @@ -1,125 +0,0 @@ -import torch -import torch.nn.functional as F -from dataclasses import dataclass -from core.reparam_module import ReparamPolicy -from core.sampling import rollout -from core.trpo import trpo_step -from core.ppo import ppo_step -from tqdm import tqdm - -class TerminalLogger: - def add_scalar(self, key, scalar, i=None): - if i is not None: - print('Iteration', i, end=' ') - print(key, scalar) - -@dataclass -class Buffer: - states: torch.Tensor - actions: torch.Tensor - rewards: torch.Tensor - dones: torch.Tensor - -def roll_buffer(buffer, *args, **kwargs): - return Buffer( - torch.roll(buffer.states, *args, **kwargs), - torch.roll(buffer.actions, *args, **kwargs), - torch.roll(buffer.rewards, *args, **kwargs), - torch.roll(buffer.dones, *args, **kwargs), - ) - -def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, - v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, - gae_lambda, delta, backtrack_coeff, backtrack_iters, cg_iters=10, cg_damping=0.1, wasserstein=False, wasserstein_c=None, logger=TerminalLogger()): - - policy(torch.zeros(env_fn(0).observation_space.shape)) - policy = ReparamPolicy(policy) - - logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0]) - logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0]) - - for epoch in tqdm(range(epochs)): - generator_data = Buffer(*rollout(env_fn, policy, rollout_episodes, rollout_steps)) - - logger.add_scalar('gen/mean_episode_length', (~generator_data.dones).sum() / generator_data.states.shape[0], epoch) - logger.add_scalar('gen/mean_reward_per_episode', generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0], epoch) - - discriminator, loss = train_discriminator(expert_data, generator_data, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c) - if wasserstein: - generator_data.rewards = discriminator(generator_data.states, generator_data.actions) - else: - generator_data.rewards = -F.logsigmoid(discriminator(generator_data.states, generator_data.actions)) - logger.add_scalar('disc/final_loss', loss, epoch) - logger.add_scalar('disc/mean_reward_per_episode', generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0], epoch) - - value, policy = trpo_step(value, policy, generator_data.states, generator_data.actions, generator_data.rewards, generator_data.dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters, cg_damping) - expert_data = roll_buffer(expert_data, shifts=-3, dims=0) - - return value, policy - -def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, - v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, - gae_lambda, clip_ratio, pi_opt, pi_iters, target_kl=None, max_grad_norm=None, wasserstein=False, wasserstein_c=None, logger=TerminalLogger()): - - logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0]) - logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0]) - - for epoch in range(epochs): - generator_data = Buffer(*rollout(env_fn, policy, rollout_episodes, rollout_steps)) - - logger.add_scalar('gen/mean_episode_length', (~generator_data.dones).sum() / generator_data.states.shape[0], epoch) - logger.add_scalar('gen/mean_reward_per_episode', generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0], epoch) - - discriminator, loss = train_discriminator(expert_data, generator_data, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c) - if wasserstein: - generator_data.rewards = discriminator(generator_data.states, generator_data.actions) - else: - generator_data.rewards = -F.logsigmoid(discriminator(generator_data.states, generator_data.actions)) - logger.add_scalar('disc/final_loss', loss, epoch) - logger.add_scalar('disc/mean_reward_per_episode', generator_data.rewards[~generator_data.dones].sum() / generator_data.states.shape[0], epoch) - - value, policy = ppo_step(value, policy, generator_data.states, generator_data.actions, generator_data.rewards, generator_data.dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm) - expert_data = roll_buffer(expert_data, shifts=-3, dims=0) - - return value, policy - -def train_discriminator(expert_data, generator_data, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c=None): - - n_expert_samples = (~expert_data.dones).sum() - n_generator_samples = (~generator_data.dones).sum() - n_samples = torch.minimum(n_expert_samples, n_generator_samples) - - gen_states = generator_data.states[~generator_data.dones][:n_samples] - gen_actions = generator_data.actions[~generator_data.dones][:n_samples] - exp_states = expert_data.states[~expert_data.dones][:n_samples] - exp_actions = expert_data.actions[~expert_data.dones][:n_samples] - - states = torch.cat((exp_states, gen_states), dim=0).detach() - actions = torch.cat((exp_actions, gen_actions), dim=0).detach() - labels = torch.cat((torch.zeros(n_samples), torch.ones(n_samples))).detach() - - # print('Batch augmentation on') - # random_states = torch.rand_like(gen_states) - # random_actions = torch.rand_like(gen_actions) - # states = torch.cat((exp_states, gen_states, random_states), dim=0).detach() - # actions = torch.cat((exp_actions, gen_actions, random_actions), dim=0).detach() - # labels = torch.cat((torch.zeros(n_samples), torch.ones(n_samples), torch.ones(n_samples))).detach() - - for _ in range(disc_iters): - disc_opt.zero_grad() - pred = discriminator(states, actions) - - if wasserstein: - loss = -(pred * (1 - labels) - pred * labels).mean() - else: - loss = F.binary_cross_entropy(torch.sigmoid(pred), labels) - - loss.backward() - disc_opt.step() - - if wasserstein_c is not None: - with torch.no_grad(): - for param in discriminator.parameters(): - param.clamp_(-wasserstein_c, wasserstein_c) - - return discriminator, loss diff --git a/scratch/etienne/trpo/core/optimization.py b/scratch/etienne/trpo/core/optimization.py deleted file mode 100644 index 4057214..0000000 --- a/scratch/etienne/trpo/core/optimization.py +++ /dev/null @@ -1,39 +0,0 @@ -import torch - -def conjugate_gradient(A, b, max_iters, res_tol=1e-10): - x = torch.zeros_like(b) - r = b - A(x) - p = r - - rTr = r.T @ r - - for _ in range(max_iters): - Ap = A(p) - alpha = rTr / (p.T @ Ap) - x = x + alpha * p - - r = r - alpha * Ap - if torch.norm(r) < res_tol: - break - - rTrnew = r.T @ r - beta = rTrnew / rTr - p = r + beta * p - rTr = rTrnew - - return x - -def line_search(f, x0, dx, g0, alpha, condition, max_steps=10, c1=0.1): - assert 0 < alpha < 1 - - f0 = f(x0) - for _ in range(max_steps): - x = x0 + dx - - if (f(x) > f0 + c1 * g0.T @ dx) and condition(x): - return x - - dx *= alpha - - print('Line search failed, returning x0') - return x0 diff --git a/scratch/etienne/trpo/core/policy.py b/scratch/etienne/trpo/core/policy.py deleted file mode 100644 index 579e72b..0000000 --- a/scratch/etienne/trpo/core/policy.py +++ /dev/null @@ -1,96 +0,0 @@ -import torch -import torch.nn as nn -from torch.distributions import Independent, Normal, Categorical -from torch.distributions.kl import kl_divergence - -class BasePolicy(nn.Module): - - def __init__(self, action_dim): - super().__init__() - self.action_dim = action_dim - - def torch_dist(self, dist): - return Independent(Normal(dist[..., :self.action_dim], dist[..., self.action_dim:].exp()), 1) - - def sample(self, dist): - return self.torch_dist(dist).sample() - - def predict(self, states): - return self.sample(self.forward(states)) - - def log_prob(self, dist, actions): - return self.torch_dist(dist).log_prob(actions) - - def kl_divergence(self, dist1, dist2): - d1 = self.torch_dist(dist1) - d2 = self.torch_dist(dist2) - return kl_divergence(d1, d2) - -class Policy(BasePolicy): - - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - self.nn = nn.Sequential( - nn.LazyLinear(50), - nn.Tanh(), - nn.LazyLinear(50), - nn.Tanh(), - nn.LazyLinear(2 * self.action_dim), - ) - - def forward(self, states): - return self.nn(states) - -class DiscretePolicy(BasePolicy): - - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - self.nn = nn.Sequential( - nn.LazyLinear(50), - nn.Tanh(), - nn.LazyLinear(50), - nn.Tanh(), - nn.LazyLinear(self.action_dim), - ) - - def forward(self, states): - return self.nn(states) - - def torch_dist(self, dist): - return Categorical(logits=dist) - -class SetPolicy(Policy): - - def forward(self, states): - batch_size = states.shape[:-2] - states = torch.cat((states[..., :1, [0, 1]], states[..., :, [2, 5]]), axis=-2).reshape(*batch_size, -1) - return super().forward(states) - -class SetDiscretePolicy(DiscretePolicy): - - def forward(self, states): - batch_size = states.shape[:-2] - states = torch.cat((states[..., :1, [0, 1]], states[..., :, [2, 5]]), axis=-2).reshape(*batch_size, -1) - return super().forward(states) - -class DeepSetPolicy(BasePolicy): - - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - self.elem = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - ) - self.glob = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(2 * self.action_dim), - ) - - def forward(self, states): - return self.glob(self.elem(states).sum(-2)) diff --git a/scratch/etienne/trpo/core/ppo.py b/scratch/etienne/trpo/core/ppo.py deleted file mode 100644 index c742061..0000000 --- a/scratch/etienne/trpo/core/ppo.py +++ /dev/null @@ -1,72 +0,0 @@ -import torch -from core.sampling import rollout -from core.value_estimation import gae - -def ppo(env_fn, value, policy, epochs, rollout_episodes, rollout_steps, gamma, gae_lambda, clip_ratio, pi_opt, pi_iters, v_opt, v_iters, target_kl=None, max_grad_norm=None): - - for epoch in range(epochs): - policy.eval() - states, actions, rewards, dones = rollout(env_fn, policy, rollout_episodes, rollout_steps) - - print('mean', states[~dones].mean(0)) - print('std', states[~dones].std(0)) - - print(f'Iteration {epoch} mean episode length {(~dones).sum() / states.shape[0]}') - print(f'Iteration {epoch} mean reward per episode {rewards[~dones].sum() / states.shape[0]}') - - policy.train() - value.train() - value, policy = ppo_step(value, policy, states, actions, rewards, dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm) - - return value, policy - -def ppo_step(value, policy, states, actions, rewards, dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm): - - states = states.detach() - actions = actions.detach() - rewards = rewards.detach() - dones = dones.detach() - - advantages, returns, valid = gae(states, rewards, value(states), dones, gamma, gae_lambda) - advantages = advantages.detach() - returns = returns.detach() - - # update value function - - for _ in range(v_iters): - v_opt.zero_grad() - value_loss = (value(states) - returns).pow(2)[valid].mean() - value_loss.backward() - v_opt.step() - - # update policy - - old_dist = policy(states).detach() - old_logprob = policy.log_prob(old_dist, actions).detach() - - def g(advantages, clip_ratio): - return torch.where(advantages >= 0, (1 + clip_ratio) * advantages, (1 - clip_ratio) * advantages) - - def L(states, actions, advantages, clip_ratio): - return torch.minimum( - (policy.log_prob(policy(states), actions) - old_logprob).exp() * advantages, - g(advantages, clip_ratio) - )[valid].mean() - - for _ in range(pi_iters): - pi_opt.zero_grad() - ppo_loss = -L(states, actions, advantages, clip_ratio) - ppo_loss.backward() - - if max_grad_norm: - torch.nn.utils.clip_grad_norm(policy.parameters(), max_grad_norm) - - pi_opt.step() - - kl = policy.kl_divergence(policy(states), old_dist)[valid].mean() - if target_kl and kl > target_kl: - break - - print('KL', kl.item()) - - return value, policy diff --git a/scratch/etienne/trpo/core/reparam_module.py b/scratch/etienne/trpo/core/reparam_module.py deleted file mode 100644 index 5bcd613..0000000 --- a/scratch/etienne/trpo/core/reparam_module.py +++ /dev/null @@ -1,162 +0,0 @@ -# Source: https://github.com/SsnL/PyTorch-Reparam-Module - -import torch -import torch.nn as nn -import warnings -import types -from collections import namedtuple -from contextlib import contextmanager - -class ReparamModule(nn.Module): - def __init__(self, module): - super(ReparamModule, self).__init__() - self.module = module - - param_infos = [] - shared_param_memo = {} - shared_param_infos = [] - params = [] - param_numels = [] - param_shapes = [] - for m in self.modules(): - for n, p in m.named_parameters(recurse=False): - if p is not None: - if p in shared_param_memo: - shared_m, shared_n = shared_param_memo[p] - shared_param_infos.append((m, n, shared_m, shared_n)) - else: - shared_param_memo[p] = (m, n) - param_infos.append((m, n)) - params.append(p.detach()) - param_numels.append(p.numel()) - param_shapes.append(p.size()) - - assert len(set(p.dtype for p in params)) <= 1, \ - "expects all parameters in module to have same dtype" - - # store the info for unflatten - self._param_infos = tuple(param_infos) - self._shared_param_infos = tuple(shared_param_infos) - self._param_numels = tuple(param_numels) - self._param_shapes = tuple(param_shapes) - - # flatten - flat_param = nn.Parameter(torch.cat([p.reshape(-1) for p in params], 0)) - self.register_parameter('flat_param', flat_param) - self.param_numel = flat_param.numel() - del params - del shared_param_memo - - # deregister the names as parameters - for m, n in self._param_infos: - delattr(m, n) - for m, n, _, _ in self._shared_param_infos: - delattr(m, n) - - # register the views as plain attributes - self._unflatten_param(self.flat_param) - - # now buffers - # they are not reparametrized. just store info as (module, name, buffer) - buffer_infos = [] - for m in self.modules(): - for n, b in m.named_buffers(recurse=False): - if b is not None: - buffer_infos.append((m, n, b)) - - self._buffer_infos = tuple(buffer_infos) - self._traced_self = None - - def trace(self, example_input, **trace_kwargs): - assert self._traced_self is None, 'This ReparamModule is already traced' - - if isinstance(example_input, torch.Tensor): - example_input = (example_input,) - example_input = tuple(example_input) - example_param = (self.flat_param.detach().clone(),) - example_buffers = (tuple(b.detach().clone() for _, _, b in self._buffer_infos),) - - self._traced_self = torch.jit.trace_module( - self, - inputs=dict( - _forward_with_param=example_param + example_input, - _forward_with_param_and_buffers=example_param + example_buffers + example_input, - ), - **trace_kwargs, - ) - - # replace forwards with traced versions - self._forward_with_param = self._traced_self._forward_with_param - self._forward_with_param_and_buffers = self._traced_self._forward_with_param_and_buffers - return self - - def clear_views(self): - for m, n in self._param_infos: - setattr(m, n, None) # This will set as plain attr - - def _apply(self, *args, **kwargs): - if self._traced_self is not None: - self._traced_self._apply(*args, **kwargs) - return self - return super(ReparamModule, self)._apply(*args, **kwargs) - - def _unflatten_param(self, flat_param): - ps = (t.view(s) for (t, s) in zip(flat_param.split(self._param_numels), self._param_shapes)) - for (m, n), p in zip(self._param_infos, ps): - setattr(m, n, p) # This will set as plain attr - for (m, n, shared_m, shared_n) in self._shared_param_infos: - setattr(m, n, getattr(shared_m, shared_n)) - - @contextmanager - def unflattened_param(self, flat_param): - saved_views = [getattr(m, n) for m, n in self._param_infos] - self._unflatten_param(flat_param) - yield - # Why not just `self._unflatten_param(self.flat_param)`? - # 1. because of https://github.com/pytorch/pytorch/issues/17583 - # 2. slightly faster since it does not require reconstruct the split+view - # graph - for (m, n), p in zip(self._param_infos, saved_views): - setattr(m, n, p) - for (m, n, shared_m, shared_n) in self._shared_param_infos: - setattr(m, n, getattr(shared_m, shared_n)) - - @contextmanager - def replaced_buffers(self, buffers): - for (m, n, _), new_b in zip(self._buffer_infos, buffers): - setattr(m, n, new_b) - yield - for m, n, old_b in self._buffer_infos: - setattr(m, n, old_b) - - def _forward_with_param_and_buffers(self, flat_param, buffers, *inputs, **kwinputs): - with self.unflattened_param(flat_param): - with self.replaced_buffers(buffers): - return self.module(*inputs, **kwinputs) - - def _forward_with_param(self, flat_param, *inputs, **kwinputs): - with self.unflattened_param(flat_param): - return self.module(*inputs, **kwinputs) - - def forward(self, *inputs, flat_param=None, buffers=None, **kwinputs): - if flat_param is None: - flat_param = self.flat_param - if buffers is None: - return self._forward_with_param(flat_param, *inputs, **kwinputs) - else: - return self._forward_with_param_and_buffers(flat_param, tuple(buffers), *inputs, **kwinputs) - - -class ReparamPolicy(ReparamModule): - - def sample(self, *args, **kwargs): - return self.module.sample(*args, **kwargs) - - def log_prob(self, *args, **kwargs): - return self.module.log_prob(*args, **kwargs) - - def kl_divergence(self, *args, **kwargs): - return self.module.kl_divergence(*args, **kwargs) - - def predict(self, *args, **kwargs): - return self.module.predict(*args, **kwargs) diff --git a/scratch/etienne/trpo/core/sampling.py b/scratch/etienne/trpo/core/sampling.py deleted file mode 100644 index 66fbbef..0000000 --- a/scratch/etienne/trpo/core/sampling.py +++ /dev/null @@ -1,73 +0,0 @@ -import torch -import gym -from stable_baselines3.common.vec_env import DummyVecEnv as VecEnv -from tqdm import tqdm - -def rollout(env_fn, policy, n_episodes, max_steps_per_episode): - env = env_fn(0) - states = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.observation_space.shape) - actions = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.action_space.shape) - rewards = torch.zeros(n_episodes, max_steps_per_episode + 1) - dones = torch.ones(n_episodes, max_steps_per_episode + 1, dtype=bool) - - env = VecEnv(list(map(lambda i: (lambda: env_fn(i)), range(n_episodes)))) - - states[:, 0] = torch.tensor(env.reset()).clone().detach() - dones[:, 0] = False - - for s in range(max_steps_per_episode): - actions[:, s] = policy.sample(policy(states[:, s])).clone().detach() - - clipped_actions = actions[:, s] - if isinstance(env.action_space, gym.spaces.Box): - clipped_actions = torch.clamp(clipped_actions, torch.from_numpy(env.action_space.low), torch.from_numpy(env.action_space.high)) - - o, r, d, _ = env.step(clipped_actions) - states[:, s + 1] = torch.tensor(o).clone().detach() - rewards[:, s] = torch.tensor(r).clone().detach() - dones[:, s + 1] = torch.tensor(d).clone().detach() - - dones = dones.cumsum(1) > 0 - - states = states[:, :max_steps_per_episode] - actions = actions[:, :max_steps_per_episode] - rewards = rewards[:, :max_steps_per_episode] - dones = dones[:, :max_steps_per_episode] - - return states, actions, rewards, dones - - -def rollout_sb3(env, policy, n_episodes, max_steps_per_episode): - states = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.observation_space.shape) - actions = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.action_space.shape) - rewards = torch.zeros(n_episodes, max_steps_per_episode + 1) - dones = torch.ones(n_episodes, max_steps_per_episode + 1, dtype=bool) - - for e in tqdm(range(n_episodes)): - states[e, 0] = torch.tensor(env.reset()).clone().detach() - dones[e, 0] = False - - for s in range(max_steps_per_episode): - action, _ = policy.predict(states[e, s]) - actions[e, s] = torch.tensor(action).clone().detach() - - clipped_actions = actions[e, s] - if isinstance(env.action_space, gym.spaces.Box): - clipped_actions = torch.clamp(clipped_actions, torch.from_numpy(env.action_space.low), torch.from_numpy(env.action_space.high)) - - o, r, d, _ = env.step(clipped_actions) - states[e, s + 1] = torch.tensor(o).clone().detach() - rewards[e, s] = torch.tensor(r).clone().detach() - dones[e, s + 1] = torch.tensor(d).clone().detach() - - if d: - break - - dones = dones.cumsum(1) > 0 - - states = states[:, :max_steps_per_episode] - actions = actions[:, :max_steps_per_episode] - rewards = rewards[:, :max_steps_per_episode] - dones = dones[:, :max_steps_per_episode] - - return states, actions, rewards, dones diff --git a/scratch/etienne/trpo/core/test_optimization.py b/scratch/etienne/trpo/core/test_optimization.py deleted file mode 100644 index aa49bb5..0000000 --- a/scratch/etienne/trpo/core/test_optimization.py +++ /dev/null @@ -1,23 +0,0 @@ -import torch -from optimization import conjugate_gradient - -def test_cg_eye(): - A = torch.eye(2) - b = torch.tensor([1., 2.]) - x1 = conjugate_gradient(lambda x: A @ x, b, 2) - x2 = torch.inverse(A) @ b - assert torch.allclose(x1, x2) - -def test_cg_eyep1(): - A = torch.eye(2) + 1 - b = torch.tensor([1., 2.]) - x1 = conjugate_gradient(lambda x: A @ x, b, 2) - x2 = torch.inverse(A) @ b - assert torch.allclose(x1, x2, atol=1e-7) - -def test_cg3(): - A = torch.tensor([[4., 2.], [2., 4.]]) - b = torch.tensor([2., 1.]) - x1 = conjugate_gradient(lambda x: A @ x, b, 100) - x2 = torch.inverse(A) @ b - assert torch.allclose(x1, x2) diff --git a/scratch/etienne/trpo/core/trpo.py b/scratch/etienne/trpo/core/trpo.py deleted file mode 100644 index d91e363..0000000 --- a/scratch/etienne/trpo/core/trpo.py +++ /dev/null @@ -1,79 +0,0 @@ -import torch -from core.reparam_module import ReparamPolicy -from core.sampling import rollout -from core.value_estimation import gae -from core.optimization import conjugate_gradient, line_search - -def trpo(env_fn, value, policy, epochs, rollout_episodes, rollout_steps, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1): - - policy(torch.zeros(env_fn(0).observation_space.shape)) - policy = ReparamPolicy(policy) - - for epoch in range(epochs): - policy.eval() - states, actions, rewards, dones = rollout(env_fn, policy, rollout_episodes, rollout_steps) - - print('mean', states[~dones].mean(0)) - print('std', states[~dones].std(0)) - - print(f'Iteration {epoch} mean episode length {(~dones).sum() / states.shape[0]}') - print(f'Iteration {epoch} mean reward per episode {rewards[~dones].sum() / states.shape[0]}') - - policy.train() - value.train() - value, policy = trpo_step(value, policy, states, actions, rewards, dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters, cg_damping) - - return value, policy - -def trpo_step(value, policy, states, actions, rewards, dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1): - - states = states.detach() - actions = actions.detach() - rewards = rewards.detach() - dones = dones.detach() - - advantages, returns, valid = gae(states, rewards, value(states), dones, gamma, gae_lambda) - advantages = advantages.detach() - returns = returns.detach() - - # update value function - - for _ in range(v_iters): - v_opt.zero_grad() - value_loss = (value(states) - returns).pow(2)[valid].mean() - value_loss.backward() - v_opt.step() - - # compute policy gradient - - plogprob = policy.log_prob(policy(states), actions) - surrogate_advantage = (plogprob * advantages)[valid].sum() / states.shape[0] - g = torch.cat(torch.autograd.grad(surrogate_advantage, policy.flat_param)).detach() - - def Hx(x): - kl = policy.kl_divergence(policy(states), policy(states).detach())[valid].mean() - dKL = torch.cat(torch.autograd.grad(kl, policy.flat_param, create_graph=True)) - H_x = torch.cat(torch.autograd.grad(dKL.T @ x, policy.flat_param)).detach() - return H_x + cg_damping * x - - x = conjugate_gradient(Hx, g, cg_iters) - npg = torch.sqrt(2 * delta / (x.T @ Hx(x))) * x - - # perform line search - - def L(theta): - rplogprob = policy.log_prob(policy(states, flat_param=theta), actions) - return ((rplogprob - plogprob.detach()).exp() * advantages)[valid].sum() / advantages.shape[0] - - condition = lambda theta: policy.kl_divergence(policy(states, flat_param=theta), policy(states))[valid].mean() < delta - - x0 = policy.flat_param - g0 = torch.cat(torch.autograd.grad(L(x0), x0)) - theta = line_search(L, x0, npg, g0, backtrack_coeff, condition, max_steps=backtrack_iters) - - # update policy parameters - - with torch.no_grad(): - policy.flat_param.copy_(theta) - - return value, policy diff --git a/scratch/etienne/trpo/core/value.py b/scratch/etienne/trpo/core/value.py deleted file mode 100644 index 2be3b78..0000000 --- a/scratch/etienne/trpo/core/value.py +++ /dev/null @@ -1,48 +0,0 @@ -import torch -import torch.nn as nn -from torch.distributions import Normal -from torch.distributions.kl import kl_divergence - -class Value(nn.Module): - - def __init__(self): - super().__init__() - self.nn = nn.Sequential( - nn.LazyLinear(50), - nn.Tanh(), - nn.LazyLinear(50), - nn.Tanh(), - nn.LazyLinear(1), - ) - - def forward(self, states): - return self.nn(states).squeeze(-1) - -class SetValue(Value): - - def forward(self, states): - batch_size = states.shape[:-2] - states = torch.cat((states[..., :1, [0, 1]], states[..., :, [2, 5]]), axis=-2).reshape(*batch_size, -1) - return super().forward(states) - -class DeepSetValue(nn.Module): - - def __init__(self): - super().__init__() - self.elem = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - ) - self.glob = nn.Sequential( - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(10), - nn.Tanh(), - nn.LazyLinear(1), - ) - - def forward(self, states): - return self.glob(self.elem(states).sum(-2)).squeeze(-1) diff --git a/scratch/etienne/trpo/core/value_estimation.py b/scratch/etienne/trpo/core/value_estimation.py deleted file mode 100644 index b558ac2..0000000 --- a/scratch/etienne/trpo/core/value_estimation.py +++ /dev/null @@ -1,40 +0,0 @@ -from operator import index -import torch - -def gae(states, rewards, values, dones, gamma, gae_lambda): - assert rewards.shape == values.shape == dones.shape - n_episodes, n_steps = rewards.shape - - valid = ~dones - valid[..., -1] = False - - td = rewards + gamma * torch.roll(values, shifts=-1, dims=1) - values - adv = td.repeat(n_steps, 1, 1).transpose(0, 1) - assert adv.shape == (n_episodes, n_steps, n_steps) - - step_start, step = torch.meshgrid(torch.arange(n_steps), torch.arange(n_steps), indexing='ij') - past = step < step_start - - # add up discounted temporal differences - discount = torch.minimum(torch.tensor(gamma).log() * (step - step_start), torch.tensor(0.)).exp() - discount = discount * ~past - discount = discount * valid.unsqueeze(1) - - adv = adv * discount - adv = adv.cumsum(2) # eq. (14) - assert adv.shape == (n_episodes, n_steps, n_steps) - - # add up discounted k-advantages - lambda_discount = torch.minimum(torch.tensor(gae_lambda).log() * (step - step_start), torch.tensor(0.)).exp() - lambda_discount = lambda_discount * ~past - lambda_discount = lambda_discount * valid.unsqueeze(1) - - adv = adv * lambda_discount - adv = adv.sum(2) / (lambda_discount.sum(2) + 1e-10) # eq. (16) - - adv = (adv - adv[valid].mean()) / adv[valid].std() - assert adv.shape == rewards.shape == values.shape - - returns = adv + values - - return adv, returns, valid diff --git a/scratch/etienne/trpo/bc-intersimple-setobs2.py b/scratch/etienne/trpo/experiments/bc-intersimple-setobs2.py similarity index 96% rename from scratch/etienne/trpo/bc-intersimple-setobs2.py rename to scratch/etienne/trpo/experiments/bc-intersimple-setobs2.py index da9c51d..04d2b90 100644 --- a/scratch/etienne/trpo/bc-intersimple-setobs2.py +++ b/scratch/etienne/trpo/experiments/bc-intersimple-setobs2.py @@ -26,7 +26,7 @@ torch.save(policy.state_dict(), 'bc-intersimple-setobs2.pt') # %% import numpy as np from core.policy import SetPolicy -from wrappers import Setobs, TransformObservation, CollisionPenaltyWrapper +from util.wrappers import Setobs, TransformObservation, CollisionPenaltyWrapper from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools diff --git a/scratch/etienne/trpo/gail-intersimple-minobs.py b/scratch/etienne/trpo/experiments/gail-intersimple-minobs.py similarity index 97% rename from scratch/etienne/trpo/gail-intersimple-minobs.py rename to scratch/etienne/trpo/experiments/gail-intersimple-minobs.py index bfeed96..62d4a2e 100644 --- a/scratch/etienne/trpo/gail-intersimple-minobs.py +++ b/scratch/etienne/trpo/experiments/gail-intersimple-minobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Minobs +from util.wrappers import CollisionPenaltyWrapper, Minobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/gail-intersimple-minobs2.py b/scratch/etienne/trpo/experiments/gail-intersimple-minobs2.py similarity index 97% rename from scratch/etienne/trpo/gail-intersimple-minobs2.py rename to scratch/etienne/trpo/experiments/gail-intersimple-minobs2.py index 0e1cf76..008b8ca 100644 --- a/scratch/etienne/trpo/gail-intersimple-minobs2.py +++ b/scratch/etienne/trpo/experiments/gail-intersimple-minobs2.py @@ -8,7 +8,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Minobs +from util.wrappers import CollisionPenaltyWrapper, Minobs import numpy as np from gym.wrappers import TransformObservation from core.reparam_module import ReparamPolicy diff --git a/scratch/etienne/trpo/gail-intersimple-normobs.py b/scratch/etienne/trpo/experiments/gail-intersimple-normobs.py similarity index 97% rename from scratch/etienne/trpo/gail-intersimple-normobs.py rename to scratch/etienne/trpo/experiments/gail-intersimple-normobs.py index 1ed0bd4..881efa0 100644 --- a/scratch/etienne/trpo/gail-intersimple-normobs.py +++ b/scratch/etienne/trpo/experiments/gail-intersimple-normobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/gail-intersimple-setobs.py b/scratch/etienne/trpo/experiments/gail-intersimple-setobs.py similarity index 97% rename from scratch/etienne/trpo/gail-intersimple-setobs.py rename to scratch/etienne/trpo/experiments/gail-intersimple-setobs.py index add36e1..7be8299 100644 --- a/scratch/etienne/trpo/gail-intersimple-setobs.py +++ b/scratch/etienne/trpo/experiments/gail-intersimple-setobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Setobs +from util.wrappers import CollisionPenaltyWrapper, Setobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/gail-intersimple-setobs2-recurrent.py b/scratch/etienne/trpo/experiments/gail-intersimple-setobs2-recurrent.py similarity index 97% rename from scratch/etienne/trpo/gail-intersimple-setobs2-recurrent.py rename to scratch/etienne/trpo/experiments/gail-intersimple-setobs2-recurrent.py index 69732e6..2130b6f 100644 --- a/scratch/etienne/trpo/gail-intersimple-setobs2-recurrent.py +++ b/scratch/etienne/trpo/experiments/gail-intersimple-setobs2-recurrent.py @@ -8,7 +8,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Setobs +from util.wrappers import CollisionPenaltyWrapper, Setobs import numpy as np from gym.wrappers import TransformObservation from core.reparam_module import ReparamPolicy diff --git a/scratch/etienne/trpo/gail-intersimple-setobs2.py b/scratch/etienne/trpo/experiments/gail-intersimple-setobs2.py similarity index 98% rename from scratch/etienne/trpo/gail-intersimple-setobs2.py rename to scratch/etienne/trpo/experiments/gail-intersimple-setobs2.py index 5a0c51b..e122ee6 100644 --- a/scratch/etienne/trpo/gail-intersimple-setobs2.py +++ b/scratch/etienne/trpo/experiments/gail-intersimple-setobs2.py @@ -8,7 +8,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Setobs +from util.wrappers import CollisionPenaltyWrapper, Setobs import numpy as np from gym.wrappers import TransformObservation from core.reparam_module import ReparamPolicy diff --git a/scratch/etienne/trpo/gail-intersimple.py b/scratch/etienne/trpo/experiments/gail-intersimple.py similarity index 96% rename from scratch/etienne/trpo/gail-intersimple.py rename to scratch/etienne/trpo/experiments/gail-intersimple.py index 4fac87e..11b050d 100644 --- a/scratch/etienne/trpo/gail-intersimple.py +++ b/scratch/etienne/trpo/experiments/gail-intersimple.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat( n_rays=5, diff --git a/scratch/etienne/trpo/gail-options-minobs.py b/scratch/etienne/trpo/experiments/gail-options-minobs.py similarity index 97% rename from scratch/etienne/trpo/gail-options-minobs.py rename to scratch/etienne/trpo/experiments/gail-options-minobs.py index ca9ef13..b738261 100644 --- a/scratch/etienne/trpo/gail-options-minobs.py +++ b/scratch/etienne/trpo/experiments/gail-options-minobs.py @@ -8,7 +8,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter diff --git a/scratch/etienne/trpo/gail-options-setobs.py b/scratch/etienne/trpo/experiments/gail-options-setobs.py similarity index 97% rename from scratch/etienne/trpo/gail-options-setobs.py rename to scratch/etienne/trpo/experiments/gail-options-setobs.py index c9ca49c..288647b 100644 --- a/scratch/etienne/trpo/gail-options-setobs.py +++ b/scratch/etienne/trpo/experiments/gail-options-setobs.py @@ -8,7 +8,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter diff --git a/scratch/etienne/trpo/gail-options-setobs2.py b/scratch/etienne/trpo/experiments/gail-options-setobs2.py similarity index 89% rename from scratch/etienne/trpo/gail-options-setobs2.py rename to scratch/etienne/trpo/experiments/gail-options-setobs2.py index 267b179..953b347 100644 --- a/scratch/etienne/trpo/gail-options-setobs2.py +++ b/scratch/etienne/trpo/experiments/gail-options-setobs2.py @@ -9,7 +9,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter @@ -56,6 +56,11 @@ expert_data = torch.load('intersimple-expert-data-setobs2.pt') expert_data = Buffer(*expert_data) # %% +def callback(epoch, value, policy): + if not epoch % 10: + torch.save(policy.state_dict(), f'gail-options-setobs2-{epoch}.pt') + torch.save(value.state_dict(), f'gail-options-setobs2-value-{epoch}.pt') + value, policy = gail( env_fn=env_fn, expert_data=expert_data, @@ -66,7 +71,7 @@ value, policy = gail( value=value, v_opt=v_opt, v_iters=1000, - epochs=200, + epochs=300, rollout_episodes=60, rollout_steps=60, gamma=0.99, @@ -75,6 +80,7 @@ value, policy = gail( backtrack_coeff=0.8, backtrack_iters=10, logger=SummaryWriter(comment='gail-options-setobs2'), + callback=callback, ) torch.save(policy.state_dict(), 'gail-options-setobs2.pt') diff --git a/scratch/etienne/trpo/gail-pendulum.py b/scratch/etienne/trpo/experiments/gail-pendulum.py similarity index 100% rename from scratch/etienne/trpo/gail-pendulum.py rename to scratch/etienne/trpo/experiments/gail-pendulum.py diff --git a/scratch/etienne/trpo/gail-ppo-intersimple-minobs.py b/scratch/etienne/trpo/experiments/gail-ppo-intersimple-minobs.py similarity index 97% rename from scratch/etienne/trpo/gail-ppo-intersimple-minobs.py rename to scratch/etienne/trpo/experiments/gail-ppo-intersimple-minobs.py index 46bcb18..21afb69 100644 --- a/scratch/etienne/trpo/gail-ppo-intersimple-minobs.py +++ b/scratch/etienne/trpo/experiments/gail-ppo-intersimple-minobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Minobs +from util.wrappers import CollisionPenaltyWrapper, Minobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/gail-ppo-intersimple-normobs.py b/scratch/etienne/trpo/experiments/gail-ppo-intersimple-normobs.py similarity index 97% rename from scratch/etienne/trpo/gail-ppo-intersimple-normobs.py rename to scratch/etienne/trpo/experiments/gail-ppo-intersimple-normobs.py index 028328a..329fb46 100644 --- a/scratch/etienne/trpo/gail-ppo-intersimple-normobs.py +++ b/scratch/etienne/trpo/experiments/gail-ppo-intersimple-normobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/gail-ppo-intersimple-setobs2.py b/scratch/etienne/trpo/experiments/gail-ppo-intersimple-setobs2.py similarity index 98% rename from scratch/etienne/trpo/gail-ppo-intersimple-setobs2.py rename to scratch/etienne/trpo/experiments/gail-ppo-intersimple-setobs2.py index 4715256..1de435d 100644 --- a/scratch/etienne/trpo/gail-ppo-intersimple-setobs2.py +++ b/scratch/etienne/trpo/experiments/gail-ppo-intersimple-setobs2.py @@ -8,7 +8,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Setobs +from util.wrappers import CollisionPenaltyWrapper, Setobs import numpy as np from gym.wrappers import TransformObservation from core.reparam_module import ReparamPolicy diff --git a/scratch/etienne/trpo/gail-ppo-intersimple.py b/scratch/etienne/trpo/experiments/gail-ppo-intersimple.py similarity index 96% rename from scratch/etienne/trpo/gail-ppo-intersimple.py rename to scratch/etienne/trpo/experiments/gail-ppo-intersimple.py index 833c028..7412b3f 100644 --- a/scratch/etienne/trpo/gail-ppo-intersimple.py +++ b/scratch/etienne/trpo/experiments/gail-ppo-intersimple.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat( n_rays=5, diff --git a/scratch/etienne/trpo/gail-ppo-options-minobs.py b/scratch/etienne/trpo/experiments/gail-ppo-options-minobs.py similarity index 97% rename from scratch/etienne/trpo/gail-ppo-options-minobs.py rename to scratch/etienne/trpo/experiments/gail-ppo-options-minobs.py index b264135..f25a9ea 100644 --- a/scratch/etienne/trpo/gail-ppo-options-minobs.py +++ b/scratch/etienne/trpo/experiments/gail-ppo-options-minobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter diff --git a/scratch/etienne/trpo/gail-ppo-options-setobs.py b/scratch/etienne/trpo/experiments/gail-ppo-options-setobs.py similarity index 97% rename from scratch/etienne/trpo/gail-ppo-options-setobs.py rename to scratch/etienne/trpo/experiments/gail-ppo-options-setobs.py index 4f2d463..8fa9339 100644 --- a/scratch/etienne/trpo/gail-ppo-options-setobs.py +++ b/scratch/etienne/trpo/experiments/gail-ppo-options-setobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter diff --git a/scratch/etienne/trpo/gail-ppo-options-setobs2.py b/scratch/etienne/trpo/experiments/gail-ppo-options-setobs2.py similarity index 89% rename from scratch/etienne/trpo/gail-ppo-options-setobs2.py rename to scratch/etienne/trpo/experiments/gail-ppo-options-setobs2.py index 7d19fc1..2ad3a46 100644 --- a/scratch/etienne/trpo/gail-ppo-options-setobs2.py +++ b/scratch/etienne/trpo/experiments/gail-ppo-options-setobs2.py @@ -8,7 +8,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter @@ -57,6 +57,11 @@ expert_data = torch.load('intersimple-expert-data-setobs2.pt') expert_data = Buffer(*expert_data) # %% +def callback(epoch, value, policy): + if not epoch % 10: + torch.save(policy.state_dict(), f'gail-ppo-options-setobs2-{epoch}.pt') + torch.save(value.state_dict(), f'gail-ppo-options-setobs2-value-{epoch}.pt') + value, policy = gail_ppo( env_fn=env_fn, expert_data=expert_data, @@ -76,6 +81,7 @@ value, policy = gail_ppo( pi_opt=pi_opt, pi_iters=100, logger=SummaryWriter(comment='gail-ppo-options-setobs2'), + callback=callback, ) torch.save(policy.state_dict(), 'gail-ppo-options-setobs2.pt') diff --git a/scratch/etienne/trpo/intersimple-expert-action-profiles.ipynb b/scratch/etienne/trpo/experiments/intersimple-expert-action-profiles.ipynb similarity index 100% rename from scratch/etienne/trpo/intersimple-expert-action-profiles.ipynb rename to scratch/etienne/trpo/experiments/intersimple-expert-action-profiles.ipynb diff --git a/scratch/etienne/trpo/intersimple-expert-rollout-minobs.py b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-minobs.py similarity index 96% rename from scratch/etienne/trpo/intersimple-expert-rollout-minobs.py rename to scratch/etienne/trpo/experiments/intersimple-expert-rollout-minobs.py index 0a00f67..9cb3802 100644 --- a/scratch/etienne/trpo/intersimple-expert-rollout-minobs.py +++ b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-minobs.py @@ -4,7 +4,7 @@ from core.sampling import rollout_sb3 from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward from intersim.expert import NormalizedIntersimpleExpert -from wrappers import CollisionPenaltyWrapper, Minobs +from util.wrappers import CollisionPenaltyWrapper, Minobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/intersimple-expert-rollout-minobs2.py b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-minobs2.py similarity index 96% rename from scratch/etienne/trpo/intersimple-expert-rollout-minobs2.py rename to scratch/etienne/trpo/experiments/intersimple-expert-rollout-minobs2.py index b65ad11..e839f74 100644 --- a/scratch/etienne/trpo/intersimple-expert-rollout-minobs2.py +++ b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-minobs2.py @@ -4,7 +4,7 @@ from core.sampling import rollout_sb3 from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward from intersim.expert import NormalizedIntersimpleExpert -from wrappers import CollisionPenaltyWrapper, Minobs +from util.wrappers import CollisionPenaltyWrapper, Minobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/intersimple-expert-rollout-normobs.py b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-normobs.py similarity index 97% rename from scratch/etienne/trpo/intersimple-expert-rollout-normobs.py rename to scratch/etienne/trpo/experiments/intersimple-expert-rollout-normobs.py index e9340ae..b839a3a 100644 --- a/scratch/etienne/trpo/intersimple-expert-rollout-normobs.py +++ b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-normobs.py @@ -4,7 +4,7 @@ from core.sampling import rollout_sb3 from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward from intersim.expert import NormalizedIntersimpleExpert -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/intersimple-expert-rollout-setobs.py b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-setobs.py similarity index 96% rename from scratch/etienne/trpo/intersimple-expert-rollout-setobs.py rename to scratch/etienne/trpo/experiments/intersimple-expert-rollout-setobs.py index 0ab531b..dcf5223 100644 --- a/scratch/etienne/trpo/intersimple-expert-rollout-setobs.py +++ b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-setobs.py @@ -4,7 +4,7 @@ from core.sampling import rollout_sb3 from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward from intersim.expert import NormalizedIntersimpleExpert -from wrappers import CollisionPenaltyWrapper, Setobs +from util.wrappers import CollisionPenaltyWrapper, Setobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/intersimple-expert-rollout-setobs2.py b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-setobs2.py similarity index 96% rename from scratch/etienne/trpo/intersimple-expert-rollout-setobs2.py rename to scratch/etienne/trpo/experiments/intersimple-expert-rollout-setobs2.py index 16ecd67..28139f2 100644 --- a/scratch/etienne/trpo/intersimple-expert-rollout-setobs2.py +++ b/scratch/etienne/trpo/experiments/intersimple-expert-rollout-setobs2.py @@ -4,7 +4,7 @@ from core.sampling import rollout_sb3 from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward from intersim.expert import NormalizedIntersimpleExpert -from wrappers import CollisionPenaltyWrapper, Setobs +from util.wrappers import CollisionPenaltyWrapper, Setobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/intersimple-expert-rollout.py b/scratch/etienne/trpo/experiments/intersimple-expert-rollout.py similarity index 94% rename from scratch/etienne/trpo/intersimple-expert-rollout.py rename to scratch/etienne/trpo/experiments/intersimple-expert-rollout.py index ac0501c..d3a7deb 100644 --- a/scratch/etienne/trpo/intersimple-expert-rollout.py +++ b/scratch/etienne/trpo/experiments/intersimple-expert-rollout.py @@ -4,7 +4,7 @@ from core.sampling import rollout_sb3 from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward from intersim.expert import NormalizedIntersimpleExpert -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper env = CollisionPenaltyWrapper(IntersimpleLidarFlat( n_rays=5, diff --git a/scratch/etienne/trpo/ppo-intersimple-minobs.py b/scratch/etienne/trpo/experiments/ppo-intersimple-minobs.py similarity index 98% rename from scratch/etienne/trpo/ppo-intersimple-minobs.py rename to scratch/etienne/trpo/experiments/ppo-intersimple-minobs.py index 89f61a6..3c648ce 100644 --- a/scratch/etienne/trpo/ppo-intersimple-minobs.py +++ b/scratch/etienne/trpo/experiments/ppo-intersimple-minobs.py @@ -9,7 +9,7 @@ import torch.optim import numpy as np from gym.wrappers import TransformObservation -from wrappers import Minobs +from util.wrappers import Minobs obs_min = np.array([ [-1000, -1000, 0, -np.pi, -1e-1, 0.], diff --git a/scratch/etienne/trpo/ppo-intersimple-minobs2.py b/scratch/etienne/trpo/experiments/ppo-intersimple-minobs2.py similarity index 98% rename from scratch/etienne/trpo/ppo-intersimple-minobs2.py rename to scratch/etienne/trpo/experiments/ppo-intersimple-minobs2.py index 5bed6b3..f119e74 100644 --- a/scratch/etienne/trpo/ppo-intersimple-minobs2.py +++ b/scratch/etienne/trpo/experiments/ppo-intersimple-minobs2.py @@ -9,7 +9,7 @@ import torch.optim import numpy as np from gym.wrappers import TransformObservation -from wrappers import Minobs +from util.wrappers import Minobs obs_min = np.array([ [-1000, -1000, 0, -np.pi, -1e-1, 0.], diff --git a/scratch/etienne/trpo/ppo-intersimple-normobs.py b/scratch/etienne/trpo/experiments/ppo-intersimple-normobs.py similarity index 100% rename from scratch/etienne/trpo/ppo-intersimple-normobs.py rename to scratch/etienne/trpo/experiments/ppo-intersimple-normobs.py diff --git a/scratch/etienne/trpo/ppo-intersimple.py b/scratch/etienne/trpo/experiments/ppo-intersimple.py similarity index 100% rename from scratch/etienne/trpo/ppo-intersimple.py rename to scratch/etienne/trpo/experiments/ppo-intersimple.py diff --git a/scratch/etienne/trpo/ppo-options-minobs.py b/scratch/etienne/trpo/experiments/ppo-options-minobs.py similarity index 95% rename from scratch/etienne/trpo/ppo-options-minobs.py rename to scratch/etienne/trpo/experiments/ppo-options-minobs.py index 38aa899..009f41d 100644 --- a/scratch/etienne/trpo/ppo-options-minobs.py +++ b/scratch/etienne/trpo/experiments/ppo-options-minobs.py @@ -9,9 +9,9 @@ from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools import numpy as np -from wrappers import CollisionPenaltyWrapper, TransformObservation +from util.wrappers import CollisionPenaltyWrapper, TransformObservation -from wrappers import Minobs +from util.wrappers import Minobs from options.options import OptionsEnv obs_min = np.array([ diff --git a/scratch/etienne/trpo/ppo-pendulum.py b/scratch/etienne/trpo/experiments/ppo-pendulum.py similarity index 100% rename from scratch/etienne/trpo/ppo-pendulum.py rename to scratch/etienne/trpo/experiments/ppo-pendulum.py diff --git a/scratch/etienne/trpo/readme.md b/scratch/etienne/trpo/experiments/readme.md similarity index 100% rename from scratch/etienne/trpo/readme.md rename to scratch/etienne/trpo/experiments/readme.md diff --git a/scratch/etienne/trpo/experiments/requirements.txt b/scratch/etienne/trpo/experiments/requirements.txt new file mode 100644 index 0000000..bd1ffb4 --- /dev/null +++ b/scratch/etienne/trpo/experiments/requirements.txt @@ -0,0 +1,3 @@ +torch +stable-baselines3 +gym diff --git a/scratch/etienne/trpo/experiments/sgail-options-setobs2.py b/scratch/etienne/trpo/experiments/sgail-options-setobs2.py new file mode 100644 index 0000000..02c8bff --- /dev/null +++ b/scratch/etienne/trpo/experiments/sgail-options-setobs2.py @@ -0,0 +1,108 @@ +# %% +import gym +from safe_options.options import gail +from core.gail import Buffer +from core.value import SetValue +from safe_options.policy import SetMaskedDiscretePolicy +from core.discriminator import DeepsetDiscriminator +import torch.optim +from intersim.envs import IntersimpleLidarFlatRandom +from intersim.envs.intersimple import speed_reward +import functools +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +import numpy as np +from safe_options.options import SafeOptionsEnv +from torch.utils.tensorboard import SummaryWriter +from core.reparam_module import ReparamPolicy + +obs_min = np.array([ + [-1000, -1000, 0, -np.pi, -1e-1, 0.], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], +]).reshape(-1) + +obs_max = np.array([ + [1000, 1000, 20, np.pi, 1e-1, 0.], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], +]).reshape(-1) + +envs = [SafeOptionsEnv(Setobs( + TransformObservation(CollisionPenaltyWrapper(IntersimpleLidarFlatRandom( + n_rays=5, + reward=functools.partial( + speed_reward, + collision_penalty=0 + ), + stop_on_collision=True, + ), collision_distance=6, collision_penalty=100), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) +), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5), (10, 5)], safe_actions_collision_method='circle', abort_unsafe_collision_method='circle') for _ in range(60)] + +env_fn = lambda i: envs[i] +policy = SetMaskedDiscretePolicy(env_fn(0).action_space.n) +value = SetValue() +v_opt = torch.optim.Adam(value.parameters(), lr=1e-4) + +discriminator = DeepsetDiscriminator() +disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-4) + +expert_data = torch.load('intersimple-expert-data-setobs2.pt') +expert_data = Buffer(*expert_data) + +# %% +def callback(epoch, value, policy): + if not epoch % 10: + torch.save(policy.state_dict(), f'sgail-options-setobs2-{epoch}.pt') + torch.save(value.state_dict(), f'sgail-options-setobs2-value-{epoch}.pt') + +value, policy = gail( + env_fn=env_fn, + expert_data=expert_data, + discriminator=discriminator, + disc_opt=disc_opt, + disc_iters=100, + policy=policy, + value=value, + v_opt=v_opt, + v_iters=1000, + epochs=300, + rollout_episodes=60, + rollout_steps=60, + gamma=0.99, + gae_lambda=0.9, + delta=0.01, + backtrack_coeff=0.8, + backtrack_iters=10, + logger=SummaryWriter(comment='sgail-options-setobs2'), + callback=callback, +) + +torch.save(policy.state_dict(), 'sgail-options-setobs2.pt') + +# %% +policy = SetMaskedDiscretePolicy(env_fn(0).action_space.n) +policy(torch.zeros(env_fn(0).observation_space['observation'].shape), torch.zeros(env_fn(0).observation_space['safe_actions'].shape)) +policy = ReparamPolicy(policy) +policy.load_state_dict(torch.load('sgail-options-setobs2.pt')) + +env = env_fn(0) +obs = env.reset() +env.render(mode='post') +for i in range(300): + action = policy.sample(policy( + torch.tensor(obs['observation'], dtype=torch.float32), + torch.tensor(obs['safe_actions'], dtype=torch.float32), + )) + obs, reward, done, _ = env.step(action, render_mode='post') + print('step', i, 'reward', reward) + if done: + break +env.close() + +# %% diff --git a/scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py b/scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py new file mode 100644 index 0000000..5365a94 --- /dev/null +++ b/scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py @@ -0,0 +1,107 @@ +# %% +import gym +from safe_options.options import gail_ppo, Buffer +from core.value import SetValue +from safe_options.policy import SetMaskedDiscretePolicy +from core.discriminator import DeepsetDiscriminator +import torch.optim +from intersim.envs import IntersimpleLidarFlatRandom +from intersim.envs.intersimple import speed_reward +import functools +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +import numpy as np +from safe_options.options import SafeOptionsEnv +from torch.utils.tensorboard import SummaryWriter + +obs_min = np.array([ + [-1000, -1000, 0, -np.pi, -1e-1, 0.], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], +]).reshape(-1) + +obs_max = np.array([ + [1000, 1000, 20, np.pi, 1e-1, 0.], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], +]).reshape(-1) + +envs = [SafeOptionsEnv(Setobs( + TransformObservation(CollisionPenaltyWrapper(IntersimpleLidarFlatRandom( + n_rays=5, + reward=functools.partial( + speed_reward, + collision_penalty=0 + ), + stop_on_collision=True, + ), collision_distance=6, collision_penalty=100), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) +), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5), (10, 5)], safe_actions_collision_method='circle', abort_unsafe_collision_method='circle') for _ in range(60)] + +env_fn = lambda i: envs[i] + +policy = SetMaskedDiscretePolicy(env_fn(0).action_space.n) +pi_opt = torch.optim.Adam(policy.parameters(), lr=3e-4) + +value = SetValue() +v_opt = torch.optim.Adam(value.parameters(), lr=1e-3) + +discriminator = DeepsetDiscriminator() +disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-4) + +expert_data = torch.load('intersimple-expert-data-setobs2.pt') +expert_data = Buffer(*expert_data) + +# %% +def callback(epoch, value, policy): + if not epoch % 10: + torch.save(policy.state_dict(), f'sgail-ppo-options-setobs2-{epoch}.pt') + torch.save(value.state_dict(), f'sgail-ppo-options-setobs2-value-{epoch}.pt') + +value, policy = gail_ppo( + env_fn=env_fn, + expert_data=expert_data, + discriminator=discriminator, + disc_opt=disc_opt, + disc_iters=100, + policy=policy, + value=value, + v_opt=v_opt, + v_iters=1000, + epochs=200, + rollout_episodes=60, + rollout_steps=60, + gamma=0.99, + gae_lambda=0.9, + clip_ratio=0.2, + pi_opt=pi_opt, + pi_iters=100, + logger=SummaryWriter(comment='sgail-ppo-options-setobs2'), + callback=callback, +) + +torch.save(policy.state_dict(), 'sgail-ppo-options-setobs2.pt') + +# %% +policy = SetMaskedDiscretePolicy(env_fn(0).action_space.n) +policy(torch.zeros(env_fn(0).observation_space['observation'].shape), torch.zeros(env_fn(0).observation_space['safe_actions'].shape)) +policy.load_state_dict(torch.load('sgail-ppo-options-setobs2.pt')) + +env = env_fn(0) +obs = env.reset() +env.render(mode='post') +for i in range(300): + action = policy.sample(policy( + torch.tensor(obs['observation'], dtype=torch.float32), + torch.tensor(obs['safe_actions'], dtype=torch.float32), + )) + obs, reward, done, _ = env.step(action, render_mode='post') + print('step', i, 'reward', reward, 'safe actions', obs['safe_actions']) + if done: + break +env.close() +# %% diff --git a/scratch/etienne/trpo/trpo-intersimple-minobs.py b/scratch/etienne/trpo/experiments/trpo-intersimple-minobs.py similarity index 96% rename from scratch/etienne/trpo/trpo-intersimple-minobs.py rename to scratch/etienne/trpo/experiments/trpo-intersimple-minobs.py index 93081d9..b3dc363 100644 --- a/scratch/etienne/trpo/trpo-intersimple-minobs.py +++ b/scratch/etienne/trpo/experiments/trpo-intersimple-minobs.py @@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward import functools import numpy as np from gym.wrappers import TransformObservation -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper from core.reparam_module import ReparamPolicy -from wrappers import Minobs +from util.wrappers import Minobs obs_min = np.array([ [-1000, -1000, 0, -np.pi, -1e-1, 0.], diff --git a/scratch/etienne/trpo/trpo-intersimple-minobs2.py b/scratch/etienne/trpo/experiments/trpo-intersimple-minobs2.py similarity index 96% rename from scratch/etienne/trpo/trpo-intersimple-minobs2.py rename to scratch/etienne/trpo/experiments/trpo-intersimple-minobs2.py index 58f6b28..0321938 100644 --- a/scratch/etienne/trpo/trpo-intersimple-minobs2.py +++ b/scratch/etienne/trpo/experiments/trpo-intersimple-minobs2.py @@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward import functools import numpy as np from gym.wrappers import TransformObservation -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper from core.reparam_module import ReparamPolicy -from wrappers import Minobs +from util.wrappers import Minobs obs_min = np.array([ [-1000, -1000, 0, -np.pi, -1e-1, 0.], diff --git a/scratch/etienne/trpo/trpo-intersimple-normobs.py b/scratch/etienne/trpo/experiments/trpo-intersimple-normobs.py similarity index 100% rename from scratch/etienne/trpo/trpo-intersimple-normobs.py rename to scratch/etienne/trpo/experiments/trpo-intersimple-normobs.py diff --git a/scratch/etienne/trpo/trpo-intersimple-setobs.py b/scratch/etienne/trpo/experiments/trpo-intersimple-setobs.py similarity index 96% rename from scratch/etienne/trpo/trpo-intersimple-setobs.py rename to scratch/etienne/trpo/experiments/trpo-intersimple-setobs.py index 62545ee..0dd77d9 100644 --- a/scratch/etienne/trpo/trpo-intersimple-setobs.py +++ b/scratch/etienne/trpo/experiments/trpo-intersimple-setobs.py @@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward import functools import numpy as np from gym.wrappers import TransformObservation -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper from core.reparam_module import ReparamPolicy -from wrappers import Setobs +from util.wrappers import Setobs obs_min = np.array([ [-1000, -1000, 0, -np.pi, -1e-1, 0.], diff --git a/scratch/etienne/trpo/trpo-intersimple-setobs2.py b/scratch/etienne/trpo/experiments/trpo-intersimple-setobs2.py similarity index 96% rename from scratch/etienne/trpo/trpo-intersimple-setobs2.py rename to scratch/etienne/trpo/experiments/trpo-intersimple-setobs2.py index 778e9d3..a64e410 100644 --- a/scratch/etienne/trpo/trpo-intersimple-setobs2.py +++ b/scratch/etienne/trpo/experiments/trpo-intersimple-setobs2.py @@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward import functools import numpy as np from gym.wrappers import TransformObservation -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper from core.reparam_module import ReparamPolicy -from wrappers import Setobs +from util.wrappers import Setobs obs_min = np.array([ [-1000, -1000, 0, -np.pi, -1e-1, 0.], diff --git a/scratch/etienne/trpo/trpo-intersimple.py b/scratch/etienne/trpo/experiments/trpo-intersimple.py similarity index 100% rename from scratch/etienne/trpo/trpo-intersimple.py rename to scratch/etienne/trpo/experiments/trpo-intersimple.py diff --git a/scratch/etienne/trpo/trpo-options-minobs.py b/scratch/etienne/trpo/experiments/trpo-options-minobs.py similarity index 95% rename from scratch/etienne/trpo/trpo-options-minobs.py rename to scratch/etienne/trpo/experiments/trpo-options-minobs.py index eaee086..dbc5a09 100644 --- a/scratch/etienne/trpo/trpo-options-minobs.py +++ b/scratch/etienne/trpo/experiments/trpo-options-minobs.py @@ -9,10 +9,10 @@ from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools import numpy as np -from wrappers import CollisionPenaltyWrapper, TransformObservation +from util.wrappers import CollisionPenaltyWrapper, TransformObservation from core.reparam_module import ReparamPolicy -from wrappers import Minobs +from util.wrappers import Minobs from options.options import OptionsEnv obs_min = np.array([ diff --git a/scratch/etienne/trpo/trpo-pendulum-rollout.py b/scratch/etienne/trpo/experiments/trpo-pendulum-rollout.py similarity index 100% rename from scratch/etienne/trpo/trpo-pendulum-rollout.py rename to scratch/etienne/trpo/experiments/trpo-pendulum-rollout.py diff --git a/scratch/etienne/trpo/trpo-pendulum.py b/scratch/etienne/trpo/experiments/trpo-pendulum.py similarity index 100% rename from scratch/etienne/trpo/trpo-pendulum.py rename to scratch/etienne/trpo/experiments/trpo-pendulum.py diff --git a/scratch/etienne/trpo/trpo-walker.py b/scratch/etienne/trpo/experiments/trpo-walker.py similarity index 100% rename from scratch/etienne/trpo/trpo-walker.py rename to scratch/etienne/trpo/experiments/trpo-walker.py diff --git a/scratch/etienne/trpo/experiments/vec-env.ipynb b/scratch/etienne/trpo/experiments/vec-env.ipynb new file mode 100644 index 0000000..1569b69 --- /dev/null +++ b/scratch/etienne/trpo/experiments/vec-env.ipynb @@ -0,0 +1,346 @@ +{ + "cells": [ + { + "cell_type": "code", + "execution_count": 3, + "metadata": {}, + "outputs": [], + "source": [ + "from stable_baselines3.common.env_util import make_vec_env\n", + "import numpy as np" + ] + }, + { + "cell_type": "code", + "execution_count": 4, + "metadata": {}, + "outputs": [], + "source": [ + "env = make_vec_env('Pendulum-v0', n_envs=6)" + ] + }, + { + "cell_type": "code", + "execution_count": 5, + "metadata": {}, + "outputs": [ + { + "data": { + "text/plain": [ + "(6, 3)" + ] + }, + "execution_count": 5, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "obs = env.reset()\n", + "obs.shape" + ] + }, + { + "cell_type": "code", + "execution_count": 6, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "1\n", + "2\n", + "3\n", + "4\n", + "5\n", + "6\n", + "7\n", + "8\n", + "9\n", + "10\n", + "11\n", + "12\n", + "13\n", + "14\n", + "15\n", + "16\n", + "17\n", + "18\n", + "19\n", + "20\n", + "21\n", + "22\n", + "23\n", + "24\n", + "25\n", + "26\n", + "27\n", + "28\n", + "29\n", + "30\n", + "31\n", + "32\n", + "33\n", + "34\n", + "35\n", + "36\n", + "37\n", + "38\n", + "39\n", + "40\n", + "41\n", + "42\n", + "43\n", + "44\n", + "45\n", + "46\n", + "47\n", + "48\n", + "49\n", + "50\n", + "51\n", + "52\n", + "53\n", + "54\n", + "55\n", + "56\n", + "57\n", + "58\n", + "59\n", + "60\n", + "61\n", + "62\n", + "63\n", + "64\n", + "65\n", + "66\n", + "67\n", + "68\n", + "69\n", + "70\n", + "71\n", + "72\n", + "73\n", + "74\n", + "75\n", + "76\n", + "77\n", + "78\n", + "79\n", + "80\n", + "81\n", + "82\n", + "83\n", + "84\n", + "85\n", + "86\n", + "87\n", + "88\n", + "89\n", + "90\n", + "91\n", + "92\n", + "93\n", + "94\n", + "95\n", + "96\n", + "97\n", + "98\n", + "99\n", + "100\n", + "101\n", + "102\n", + "103\n", + "104\n", + "105\n", + "106\n", + "107\n", + "108\n", + "109\n", + "110\n", + "111\n", + "112\n", + "113\n", + "114\n", + "115\n", + "116\n", + "117\n", + "118\n", + "119\n", + "120\n", + "121\n", + "122\n", + "123\n", + "124\n", + "125\n", + "126\n", + "127\n", + "128\n", + "129\n", + "130\n", + "131\n", + "132\n", + "133\n", + "134\n", + "135\n", + "136\n", + "137\n", + "138\n", + "139\n", + "140\n", + "141\n", + "142\n", + "143\n", + "144\n", + "145\n", + "146\n", + "147\n", + "148\n", + "149\n", + "150\n", + "151\n", + "152\n", + "153\n", + "154\n", + "155\n", + "156\n", + "157\n", + "158\n", + "159\n", + "160\n", + "161\n", + "162\n", + "163\n", + "164\n", + "165\n", + "166\n", + "167\n", + "168\n", + "169\n", + "170\n", + "171\n", + "172\n", + "173\n", + "174\n", + "175\n", + "176\n", + "177\n", + "178\n", + "179\n", + "180\n", + "181\n", + "182\n", + "183\n", + "184\n", + "185\n", + "186\n", + "187\n", + "188\n", + "189\n", + "190\n", + "191\n", + "192\n", + "193\n", + "194\n", + "195\n", + "196\n", + "197\n", + "198\n", + "199\n", + "200\n" + ] + } + ], + "source": [ + "dones = [False]\n", + "i = 0\n", + "while not any(dones):\n", + " i += 1\n", + " print(i)\n", + " _, _, dones, _ = env.step(np.zeros((6, 1)))" + ] + }, + { + "cell_type": "code", + "execution_count": 7, + "metadata": {}, + "outputs": [ + { + "data": { + "text/plain": [ + "array([ True, True, True, True, True, True])" + ] + }, + "execution_count": 7, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "dones" + ] + }, + { + "cell_type": "code", + "execution_count": 8, + "metadata": {}, + "outputs": [], + "source": [ + "_, _, dones, _ = env.step(np.zeros((6, 1)))" + ] + }, + { + "cell_type": "code", + "execution_count": 9, + "metadata": {}, + "outputs": [ + { + "data": { + "text/plain": [ + "array([False, False, False, False, False, False])" + ] + }, + "execution_count": 9, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "dones" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [] + } + ], + "metadata": { + "interpreter": { + "hash": "6c7a4ac80dd345f83235e10baa3acc437d966916e1cc075a45b91bb9cc030938" + }, + "kernelspec": { + "display_name": "Python 3.9.7 64-bit ('.venv': venv)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.9.7" + }, + "orig_nbformat": 4 + }, + "nbformat": 4, + "nbformat_minor": 2 +} diff --git a/scratch/etienne/trpo/wgail-intersimple-minobs.py b/scratch/etienne/trpo/experiments/wgail-intersimple-minobs.py similarity index 97% rename from scratch/etienne/trpo/wgail-intersimple-minobs.py rename to scratch/etienne/trpo/experiments/wgail-intersimple-minobs.py index b36988e..677cfef 100644 --- a/scratch/etienne/trpo/wgail-intersimple-minobs.py +++ b/scratch/etienne/trpo/experiments/wgail-intersimple-minobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Minobs +from util.wrappers import CollisionPenaltyWrapper, Minobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/wgail-intersimple-minobs2.py b/scratch/etienne/trpo/experiments/wgail-intersimple-minobs2.py similarity index 97% rename from scratch/etienne/trpo/wgail-intersimple-minobs2.py rename to scratch/etienne/trpo/experiments/wgail-intersimple-minobs2.py index 733a95c..1b97041 100644 --- a/scratch/etienne/trpo/wgail-intersimple-minobs2.py +++ b/scratch/etienne/trpo/experiments/wgail-intersimple-minobs2.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Minobs +from util.wrappers import CollisionPenaltyWrapper, Minobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/wgail-intersimple-setobs2.py b/scratch/etienne/trpo/experiments/wgail-intersimple-setobs2.py similarity index 97% rename from scratch/etienne/trpo/wgail-intersimple-setobs2.py rename to scratch/etienne/trpo/experiments/wgail-intersimple-setobs2.py index 95b8ebc..08e653c 100644 --- a/scratch/etienne/trpo/wgail-intersimple-setobs2.py +++ b/scratch/etienne/trpo/experiments/wgail-intersimple-setobs2.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Setobs +from util.wrappers import CollisionPenaltyWrapper, Setobs import numpy as np from gym.wrappers import TransformObservation from torch.utils.tensorboard import SummaryWriter diff --git a/scratch/etienne/trpo/wgail-intersimple.py b/scratch/etienne/trpo/experiments/wgail-intersimple.py similarity index 96% rename from scratch/etienne/trpo/wgail-intersimple.py rename to scratch/etienne/trpo/experiments/wgail-intersimple.py index 793cb60..57d1f40 100644 --- a/scratch/etienne/trpo/wgail-intersimple.py +++ b/scratch/etienne/trpo/experiments/wgail-intersimple.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat( n_rays=5, diff --git a/scratch/etienne/trpo/wgail-options-setobs.py b/scratch/etienne/trpo/experiments/wgail-options-setobs.py similarity index 97% rename from scratch/etienne/trpo/wgail-options-setobs.py rename to scratch/etienne/trpo/experiments/wgail-options-setobs.py index 89944ed..eaf2d5b 100644 --- a/scratch/etienne/trpo/wgail-options-setobs.py +++ b/scratch/etienne/trpo/experiments/wgail-options-setobs.py @@ -9,7 +9,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter diff --git a/scratch/etienne/trpo/wgail-options-setobs2.py b/scratch/etienne/trpo/experiments/wgail-options-setobs2.py similarity index 96% rename from scratch/etienne/trpo/wgail-options-setobs2.py rename to scratch/etienne/trpo/experiments/wgail-options-setobs2.py index 6f7b2af..a851a46 100644 --- a/scratch/etienne/trpo/wgail-options-setobs2.py +++ b/scratch/etienne/trpo/experiments/wgail-options-setobs2.py @@ -9,7 +9,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter @@ -50,7 +50,7 @@ value = SetValue() v_opt = torch.optim.Adam(value.parameters(), lr=1e-4) discriminator = DeepsetDiscriminator() -disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-3) +disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-4) expert_data = torch.load('intersimple-expert-data-setobs2.pt') expert_data = Buffer(*expert_data) diff --git a/scratch/etienne/trpo/wgail-pendulum.py b/scratch/etienne/trpo/experiments/wgail-pendulum.py similarity index 100% rename from scratch/etienne/trpo/wgail-pendulum.py rename to scratch/etienne/trpo/experiments/wgail-pendulum.py diff --git a/scratch/etienne/trpo/wgail-ppo-intersimple-minobs.py b/scratch/etienne/trpo/experiments/wgail-ppo-intersimple-minobs.py similarity index 97% rename from scratch/etienne/trpo/wgail-ppo-intersimple-minobs.py rename to scratch/etienne/trpo/experiments/wgail-ppo-intersimple-minobs.py index 45d0143..ee9757d 100644 --- a/scratch/etienne/trpo/wgail-ppo-intersimple-minobs.py +++ b/scratch/etienne/trpo/experiments/wgail-ppo-intersimple-minobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Minobs +from util.wrappers import CollisionPenaltyWrapper, Minobs import numpy as np from gym.wrappers import TransformObservation diff --git a/scratch/etienne/trpo/wgail-ppo-intersimple-setobs2.py b/scratch/etienne/trpo/experiments/wgail-ppo-intersimple-setobs2.py similarity index 97% rename from scratch/etienne/trpo/wgail-ppo-intersimple-setobs2.py rename to scratch/etienne/trpo/experiments/wgail-ppo-intersimple-setobs2.py index 30b69a0..fe7472e 100644 --- a/scratch/etienne/trpo/wgail-ppo-intersimple-setobs2.py +++ b/scratch/etienne/trpo/experiments/wgail-ppo-intersimple-setobs2.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, Setobs +from util.wrappers import CollisionPenaltyWrapper, Setobs import numpy as np from gym.wrappers import TransformObservation from torch.utils.tensorboard import SummaryWriter diff --git a/scratch/etienne/trpo/wgail-ppo-intersimple.py b/scratch/etienne/trpo/experiments/wgail-ppo-intersimple.py similarity index 100% rename from scratch/etienne/trpo/wgail-ppo-intersimple.py rename to scratch/etienne/trpo/experiments/wgail-ppo-intersimple.py diff --git a/scratch/etienne/trpo/wgail-ppo-options-setobs.py b/scratch/etienne/trpo/experiments/wgail-ppo-options-setobs.py similarity index 97% rename from scratch/etienne/trpo/wgail-ppo-options-setobs.py rename to scratch/etienne/trpo/experiments/wgail-ppo-options-setobs.py index fa89bf0..3f7ed6d 100644 --- a/scratch/etienne/trpo/wgail-ppo-options-setobs.py +++ b/scratch/etienne/trpo/experiments/wgail-ppo-options-setobs.py @@ -7,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter diff --git a/scratch/etienne/trpo/wgail-ppo-options-setobs2.py b/scratch/etienne/trpo/experiments/wgail-ppo-options-setobs2.py similarity index 96% rename from scratch/etienne/trpo/wgail-ppo-options-setobs2.py rename to scratch/etienne/trpo/experiments/wgail-ppo-options-setobs2.py index ea6fbfa..d4394a2 100644 --- a/scratch/etienne/trpo/wgail-ppo-options-setobs2.py +++ b/scratch/etienne/trpo/experiments/wgail-ppo-options-setobs2.py @@ -1,4 +1,3 @@ -# %% import gym from options.options import gail_ppo, Buffer from core.value import SetValue @@ -8,7 +7,7 @@ import torch.optim from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs.intersimple import speed_reward import functools -from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs +from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np from options.options import OptionsEnv from torch.utils.tensorboard import SummaryWriter @@ -51,12 +50,11 @@ value = SetValue() v_opt = torch.optim.Adam(value.parameters(), lr=1e-3) discriminator = DeepsetDiscriminator() -disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-3) +disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-4) expert_data = torch.load('intersimple-expert-data-setobs2.pt') expert_data = Buffer(*expert_data) -# %% value, policy = gail_ppo( env_fn=env_fn, expert_data=expert_data, diff --git a/scratch/etienne/trpo/wgail-ppo-pendulum.py b/scratch/etienne/trpo/experiments/wgail-ppo-pendulum.py similarity index 100% rename from scratch/etienne/trpo/wgail-ppo-pendulum.py rename to scratch/etienne/trpo/experiments/wgail-ppo-pendulum.py diff --git a/scratch/etienne/trpo/sb3/sb3-ppo-intersimple-rollout.py b/scratch/etienne/trpo/sb3/sb3-ppo-intersimple-rollout.py index eee73fe..2f9b6fd 100644 --- a/scratch/etienne/trpo/sb3/sb3-ppo-intersimple-rollout.py +++ b/scratch/etienne/trpo/sb3/sb3-ppo-intersimple-rollout.py @@ -7,7 +7,7 @@ from intersim.envs import IntersimpleLidarFlat from intersim.envs.intersimple import speed_reward import functools import torch -from wrappers import CollisionPenaltyWrapper +from util.wrappers import CollisionPenaltyWrapper model = PPO.load('sb3-ppo-intersimple') env = CollisionPenaltyWrapper(IntersimpleLidarFlat( diff --git a/src/core/reparam_module.py b/src/core/reparam_module.py index 5bcd613..1b24986 100644 --- a/src/core/reparam_module.py +++ b/src/core/reparam_module.py @@ -160,3 +160,6 @@ class ReparamPolicy(ReparamModule): def predict(self, *args, **kwargs): return self.module.predict(*args, **kwargs) + + def unsafe_probability_mass(self, *args, **kwargs): + return self.module.unsafe_probability_mass(*args, **kwargs) diff --git a/src/options/envs.py b/src/options/envs.py new file mode 100644 index 0000000..6f6bdd4 --- /dev/null +++ b/src/options/envs.py @@ -0,0 +1,111 @@ +import gym +import numpy as np +from src.gail2.wrappers import Wrapper, Setobs, TransformObservation +from intersim.envs import IntersimpleLidarFlatIncrementingAgent + +obs_min = np.array([ + [-1000, -1000, 0, -np.pi, -1e-1, 0.], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], + [0, -np.pi, -20, -20, -np.pi, -1e-1], +]).reshape(-1) + +obs_max = np.array([ + [1000, 1000, 20, np.pi, 1e-1, 0.], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], + [50, np.pi, 20, 20, np.pi, 1e-1], +]).reshape(-1) + +def NormalizedOptionsEvalEnv(**kwargs): + return OptionsEnv(Setobs( + TransformObservation(IntersimpleLidarFlatIncrementingAgent( + n_rays=5, + **kwargs, + ), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) + ), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5)]) + +def NormalizedContinuousEvalEnv(**kwargs): + return Setobs( + TransformObservation(IntersimpleLidarFlatIncrementingAgent( + n_rays=5, + **kwargs, + ), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) + ) + +class OptionsEnv(Wrapper): + + def __init__(self, env, options): + super().__init__(env) + self.ll_action_space = env.action_space + self.options = options + self.action_space = gym.spaces.Discrete(len(options)) + self.max_plan_length = max(t for _, t in options) + + def plan(self, option): + target_v, t = option + current_v = self.env._env.state[self.env._agent, 1].item() + dt = self.env._env._dt + a = (target_v - current_v) / (t * dt) + a = self.env._normalize(a) + a = a * np.ones((t,)) + a += 0.01 * np.random.randn(*a.shape) + a = np.clip(a, self.ll_action_space.low, self.ll_action_space.high) + return a + + def execute_plan(self, obs, option, render_mode=None): + observations = np.zeros((self.max_plan_length + 1, *self.env.observation_space.shape)) + actions = np.zeros((self.max_plan_length + 1, *self.ll_action_space.shape)) + rewards = np.zeros((self.max_plan_length + 1,)) + env_done = np.ones((self.max_plan_length + 1,), dtype=bool) + plan_done = np.ones((self.max_plan_length + 1,), dtype=bool) + infos = [] + + observations[0] = obs + env_done[0] = False + for k, u in enumerate(self.plan(option)): + plan_done[k] = False + o, r, d, i = super().step(u) + actions[k] = u + rewards[k] = r + env_done[k] = d + infos.append(i) + observations[k+1] = o + + if render_mode is not None: + self.env.render(render_mode) + + if d: + break + + n_steps = k + 1 + return observations, actions, rewards, env_done, plan_done, infos, n_steps + + def step(self, action, render_mode=None): + a = int(action) + assert a == action + ll_obs, ll_actions, ll_rewards, ll_env_done, ll_plan_done, ll_infos, ll_steps = self.execute_plan(self.last_obs, self.options[a], render_mode) + hl_obs = ll_obs[ll_steps] + hl_reward = (ll_rewards * ~ll_plan_done).sum().item() + hl_done = ll_env_done[ll_steps-1].item() + hl_infos = { + 'll': { + 'observations': ll_obs, + 'actions': ll_actions, + 'rewards': ll_rewards, + 'env_done': ll_env_done, + 'plan_done': ll_plan_done, + 'infos': ll_infos, + 'steps': ll_steps, + } + } + self.last_obs = hl_obs + return hl_obs, hl_reward, hl_done, hl_infos + + def reset(self, *args, **kwargs): + self.last_obs = super().reset(*args, **kwargs) + return self.last_obs diff --git a/scratch/etienne/trpo/options/options.py b/src/options/options.py similarity index 96% rename from scratch/etienne/trpo/options/options.py rename to src/options/options.py index af0414a..0453fd5 100644 --- a/scratch/etienne/trpo/options/options.py +++ b/src/options/options.py @@ -18,7 +18,7 @@ class OptionsRollout: def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, - gae_lambda, delta, backtrack_coeff, backtrack_iters, cg_iters=10, cg_damping=0.1, wasserstein=False, wasserstein_c=None, logger=TerminalLogger()): + gae_lambda, delta, backtrack_coeff, backtrack_iters, cg_iters=10, cg_damping=0.1, wasserstein=False, wasserstein_c=None, logger=TerminalLogger(), callback=None): policy(torch.zeros(env_fn(0).observation_space.shape)) policy = ReparamPolicy(policy) @@ -49,11 +49,14 @@ def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value value, policy = trpo_step(value, policy, generator_data.hl.states, generator_data.hl.actions, generator_data.hl.rewards, generator_data.hl.dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters, cg_damping) expert_data = roll_buffer(expert_data, shifts=-3, dims=0) + if callback is not None: + callback(epoch, value, policy) + return value, policy def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, - gae_lambda, clip_ratio, pi_opt, pi_iters, target_kl=None, max_grad_norm=None, wasserstein=False, wasserstein_c=None, logger=TerminalLogger()): + gae_lambda, clip_ratio, pi_opt, pi_iters, target_kl=None, max_grad_norm=None, wasserstein=False, wasserstein_c=None, logger=TerminalLogger(), callback=None): logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0]) logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0]) @@ -81,6 +84,9 @@ def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, v value, policy = ppo_step(value, policy, generator_data.hl.states, generator_data.hl.actions, generator_data.hl.rewards, generator_data.hl.dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm) expert_data = roll_buffer(expert_data, shifts=-3, dims=0) + if callback is not None: + callback(epoch, value, policy) + return value, policy def rollout(env_fn, policy, n_episodes, max_steps_per_episode): @@ -99,7 +105,6 @@ def rollout(env_fn, policy, n_episodes, max_steps_per_episode): env = VecEnv(list(map(lambda i: (lambda: env_fn(i)), range(n_episodes)))) states[:, 0] = torch.tensor(env.reset()).clone().detach() - dones[:, 0] = False for s in tqdm(range(max_steps_per_episode), 'Rollout'): actions[:, s] = policy.sample(policy(states[:, s])).clone().detach() @@ -111,7 +116,7 @@ def rollout(env_fn, policy, n_episodes, max_steps_per_episode): o, r, d, info = env.step(clipped_actions) states[:, s + 1] = torch.tensor(o).clone().detach() rewards[:, s] = torch.tensor(r).clone().detach() - dones[:, s + 1] = torch.tensor(d).clone().detach() + dones[:, s] = torch.tensor(d).clone().detach() ll_states[:, s] = torch.from_numpy(np.stack([i['ll']['observations'] for i in info])).clone().detach() ll_actions[:, s] = torch.from_numpy(np.stack([i['ll']['actions'] for i in info])).clone().detach() @@ -162,7 +167,7 @@ class OptionsEnv(gym.Wrapper): o, r, d, i = super().step(u) actions[k] = u rewards[k] = r - env_done[k+1] = d + env_done[k] = d infos.append(i) observations[k+1] = o @@ -181,7 +186,7 @@ class OptionsEnv(gym.Wrapper): ll_obs, ll_actions, ll_rewards, ll_env_done, ll_plan_done, ll_infos, ll_steps = self.execute_plan(self.last_obs, self.options[a], render_mode) hl_obs = ll_obs[ll_steps] hl_reward = (ll_rewards * ~ll_plan_done).sum().item() - hl_done = ll_env_done[ll_steps].item() + hl_done = ll_env_done[ll_steps-1].item() hl_infos = { 'll': { 'observations': ll_obs, diff --git a/scratch/etienne/trpo/options/test_options.py b/src/options/test_options.py similarity index 100% rename from scratch/etienne/trpo/options/test_options.py rename to src/options/test_options.py diff --git a/src/safe_options/collisions.py b/src/safe_options/collisions.py new file mode 100644 index 0000000..36a6448 --- /dev/null +++ b/src/safe_options/collisions.py @@ -0,0 +1,185 @@ +import torch +import numpy as np +from intersim.collisions import state_to_polygon + +def safety_plan(env, plan): + return np.concatenate((plan, np.array(5 * [env._env._min_acc])), axis=0) + +def available_actions(env, options): + """Return mask of available actions given current `env` state.""" + plans = [generate_plan(env, i, options) for i, _ in enumerate(options)] + # is emergency braking still possible? + plans = list(map(lambda p: safety_plan(env, p), plans)) + + T = max(len(p) for p in plans) + plans = [np.pad(p, ((0, T-len(p)),), constant_values=np.nan) for p in plans] + plans = np.stack(plans, axis=0) + + valid = feasible(env, plans) + return valid + +def target_velocity_plan(current_v: float, target_v: float, t: int, dt: float): + """Smoothly target a velocity in a given number of steps""" + # for now, constant acceleration + a = (target_v - current_v) / (t * dt) + return a*np.ones((t,)) + +def generate_plan(env, i, options): + """Generate input profile for high-level action `i`.""" + assert i < len(options), "Invalid option index {i}" + target_v, t = options[i] + current_v = env._env.state[env._agent, 1].item() # extract from env + plan = target_velocity_plan(current_v, target_v, t, env._env._dt) + assert len(plan) == t, "incorrect plan length" + return plan + +def feasible(env, plan, method='exact'): + """Check if input profile is feasible given current `env` state.""" + # zero pad plan - Take (B, T) or (T,) np plan and convert it to (B, T, nv, 1) torch.Tensor + plan = torch.tensor(plan) + plan = plan.reshape(-1, plan.shape[-1]) + full_plan = torch.zeros(*plan.shape, env._env._nv, 1) + full_plan[:, :, env._agent, 0] = plan + + # check_future_collisions_fast takes in B-list and outputs (B,) bool tensor + if method=='circle': + valid = check_future_collisions_fast(env, full_plan) + elif method=='ncircles': + valid = check_future_collisions_ncircles(env, full_plan) + elif method=='exact': + valid = check_future_collisions_exact(env, full_plan) + else: + raise NotImplementedError('Invalid collision-checking method') + + return valid + +def check_future_collisions_ncircles(env, actions, n_circles:int=2): + """Checks whether `env._agent` would collide with other agents assuming `actions` as input. + + Vehicles are (over-)approximated by multiple circles. + + Args: + env (gym.Env): current environment state + actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles + Returns: + feasible (torch.Tensor): tensor of shape (B,) indicating whether the respective action profiles are collision-free + """ + assert n_circles >= 2 + B, (T, nv, _) = len(actions), actions[0].shape + + states = env._env.propagate_action_profile_vectorized(actions) + assert states.shape == (B, T, nv, 5) + centers = states[:, :, :, :2] + psi = states[:, :, :, 3] + lon = torch.stack([psi.cos(), psi.sin()],dim=-1) # (B, T, nv, 2) + + # offset between [-env._env.lengths+env._env.widths/2, env._env.lengths/2-env._env.widths/2] + back = (-env._env._lengths/2+env._env._widths/2).unsqueeze(-1) # (nv, 1) + length = (env._env._lengths-env._env._widths).unsqueeze(-1) # (nv, 1) + diff_d = back + length*(torch.arange(n_circles)/(n_circles-1)).unsqueeze(0) # (nv, n_circles) + assert diff_d.shape == (nv, n_circles) + + offsets = diff_d[None, None, :, :, None] * lon[:, :, :, None, :] + assert offsets.shape == (B, T, nv, n_circles, 2) + + expanded_centers=centers.unsqueeze(-2) + offsets #(B, T, nv, n_circles, 2) + assert expanded_centers.shape == (B, T, nv, n_circles, 2) + agent_centers = expanded_centers[:,:,env._agent:env._agent+1,:,:] #(B, T, 1, n_circles, 2) + ds = expanded_centers.reshape((B, T, nv*n_circles, 1, 2)) - agent_centers #(B, T, nv*nc,1, 2) - (B, T, 1, nc, 2) = (B, T, nv*nc, nc, 2) + + distance = (ds**2).sum(-1).sqrt().reshape((B, T, nv, n_circles, n_circles)) # (B, T, nv, nc, nc) + distance = torch.where(distance.isnan(), np.inf*torch.ones_like(distance), distance) # only collide with spawned agents + distance[:, :, env._agent] = np.inf # cannot collide with itself + assert distance.shape == (B, T, nv, n_circles, n_circles) + + radius = env._env._widths*np.sqrt(2) / 2 + min_distance = radius[env._agent] + radius + min_distance = min_distance[None, None, :, None, None] + assert min_distance.shape == (1, 1, nv, 1, 1) + + return (distance > min_distance).all(-1).all(-1).all(-1).all(-1) + +def check_future_collisions_circle(env, actions): + """Compute collision information for circular vehicle approximations + + Args: + env (gym.Env): current environment state + actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles + Returns: + states (torch.Tensor): tensor of shape (B, T, nv, 5) of future states based on the action profiles + collision_tensor (torch.Tensor): tensor of shape (B, T, nv) of bools indicating which plan collides with which vehicles in which time frame + false: colliding, true: not colliding + """ + B, (T, nv, _) = len(actions), actions[0].shape + + states = env._env.propagate_action_profile_vectorized(actions) + assert states.shape == (B, T, nv, 5) + + distance = ((states[:, :, :, :2] - states[:, :, env._agent:env._agent+1, :2])**2).sum(-1).sqrt() + distance = torch.where(distance.isnan(), np.inf*torch.ones_like(distance), distance) # only collide with spawned agents + distance[:, :, env._agent] = np.inf # cannot collide with itself + assert distance.shape == (B, T, nv) + + radius = (env._env._lengths**2 + env._env._widths**2).sqrt() / 2 + min_distance = radius[env._agent] + radius + min_distance = min_distance.unsqueeze(0).unsqueeze(0) + assert min_distance.shape == (1, 1, nv) + + collision_tensor = distance > min_distance + assert collision_tensor.shape == (B, T, nv) + return states, collision_tensor + +def check_future_collisions_fast(env, actions): + """Checks whether `env._agent` would collide with other agents assuming `actions` as input. + + Vehicles are (over-)approximated by single circles. + + Args: + env (gym.Env): current environment state + actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles + Returns: + feasible (torch.Tensor): tensor of shape (B,) indicating whether the respective action profiles are collision-free + """ + _, collision_tensor = check_future_collisions_circle(env, actions) + return collision_tensor.all(-1).all(-1) + +def check_future_collisions_exact(env, actions): + """ + Checks whether `env._agent` would collide with other agents assuming `actions` as input. + + Args: + env (gym.Env): current environment state + actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles + Returns: + feasible (torch.Tensor): tensor of shape (B,) indicating whether the respective action profiles are collision-free + """ + # First check with simple circle collision check + states, collision_tensor = check_future_collisions_circle(env, actions) + (B, T, nv, _) = states.shape + # For those that have colliding circles, check exactly + colliding_mask = ~collision_tensor + + ego_states = states[:, :, env._agent:env._agent+1, :].expand(states.shape) + assert ego_states.shape == states.shape + + # get dimensions + lengths = env._env._lengths.expand(states.shape[:3]) + widths = env._env._widths.expand(states.shape[:3]) + ego_lengths = lengths[:, :, env._agent:env._agent+1].expand(lengths.shape) + ego_widths = widths[:, :, env._agent:env._agent+1].expand(widths.shape) + assert lengths.shape == widths.shape == ego_lengths.shape == ego_widths.shape == (B, T, nv) + + # For every collision instance between ego and other vehicle, check whether rectangles intersect + exact_collisions = torch.zeros_like(collision_tensor[colliding_mask]) + for i, (ego_state, ego_length, ego_width, other_state, other_length, other_width) in enumerate(zip( + ego_states[colliding_mask], ego_lengths[colliding_mask], ego_widths[colliding_mask], + states[colliding_mask], lengths[colliding_mask], widths[colliding_mask] + )): + assert ego_state.shape == other_state.shape == (5,) + assert ego_length.shape == ego_width.shape == other_length.shape == other_width.shape == () + p_ego = state_to_polygon(ego_state, ego_length, ego_width) + p_other = state_to_polygon(other_state, other_length, other_width) + exact_collisions[i] = p_ego.intersects(p_other) + + collision_tensor[colliding_mask] = ~exact_collisions + return collision_tensor.all(-1).all(-1) diff --git a/src/safe_options/options.py b/src/safe_options/options.py new file mode 100644 index 0000000..a8d0e50 --- /dev/null +++ b/src/safe_options/options.py @@ -0,0 +1,305 @@ +import gym +import numpy as np +import torch +from stable_baselines3.common.vec_env import DummyVecEnv as VecEnv + +from core.reparam_module import ReparamPolicy +from tqdm import tqdm +from core.gail import train_discriminator, roll_buffer, TerminalLogger +from dataclasses import dataclass +from safe_options.policy_gradient import trpo_step, ppo_step +import torch.nn.functional as F + +from safe_options.collisions import feasible + +@dataclass +class Buffer: + states: torch.Tensor + actions: torch.Tensor + rewards: torch.Tensor + dones: torch.Tensor + +@dataclass +class HLBuffer: + states: torch.Tensor + safe_actions: torch.Tensor + actions: torch.Tensor + rewards: torch.Tensor + dones: torch.Tensor + +@dataclass +class OptionsRollout: + hl: HLBuffer + ll: Buffer + +def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, + v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, + gae_lambda, delta, backtrack_coeff, backtrack_iters, cg_iters=10, cg_damping=0.1, wasserstein=False, wasserstein_c=None, logger=TerminalLogger(), callback=None): + + policy(torch.zeros(env_fn(0).observation_space['observation'].shape), torch.zeros(env_fn(0).observation_space['safe_actions'].shape)) + policy = ReparamPolicy(policy) + + logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0]) + logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0]) + + for epoch in tqdm(range(epochs)): + hl_data, ll_data = rollout(env_fn, policy, rollout_episodes, rollout_steps) + generator_data = OptionsRollout(HLBuffer(*hl_data), Buffer(*ll_data)) + + generator_data.ll.actions += 0.1 * torch.randn_like(generator_data.ll.actions) + + logger.add_scalar('gen/mean_episode_length', (~generator_data.ll.dones).sum() / generator_data.ll.states.shape[0], epoch) + logger.add_scalar('gen/mean_reward_per_episode', generator_data.hl.rewards[~generator_data.hl.dones].sum() / generator_data.hl.states.shape[0], epoch) + logger.add_scalar('gen/unsafe_probability_mass', policy.unsafe_probability_mass(policy(generator_data.hl.states[~generator_data.hl.dones], generator_data.hl.safe_actions[~generator_data.hl.dones])).mean(), epoch) + + discriminator, loss = train_discriminator(expert_data, generator_data.ll, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c) + if wasserstein: + generator_data.ll.rewards = discriminator(generator_data.ll.states, generator_data.ll.actions) + else: + generator_data.ll.rewards = -F.logsigmoid(discriminator(generator_data.ll.states, generator_data.ll.actions)) + logger.add_scalar('disc/final_loss', loss, epoch) + logger.add_scalar('disc/mean_reward_per_episode', generator_data.ll.rewards[~generator_data.ll.dones].sum() / generator_data.ll.states.shape[0], epoch) + + #assert generator_data.ll.rewards.shape == generator_data.ll.dones.shape + generator_data.hl.rewards = torch.where(~generator_data.ll.dones, generator_data.ll.rewards, torch.tensor(0.)).sum(-1) + + value, policy = trpo_step(value, policy, generator_data.hl.states, generator_data.hl.safe_actions, generator_data.hl.actions, generator_data.hl.rewards, generator_data.hl.dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters, cg_damping) + expert_data = roll_buffer(expert_data, shifts=-3, dims=0) + + if callback is not None: + callback(epoch, value, policy) + + return value, policy + +def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, + v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, + gae_lambda, clip_ratio, pi_opt, pi_iters, target_kl=None, max_grad_norm=None, wasserstein=False, wasserstein_c=None, logger=TerminalLogger(), callback=None): + + logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0]) + logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0]) + + for epoch in range(epochs): + hl_data, ll_data = rollout(env_fn, policy, rollout_episodes, rollout_steps) + generator_data = OptionsRollout(HLBuffer(*hl_data), Buffer(*ll_data)) + + generator_data.ll.actions += 0.1 * torch.randn_like(generator_data.ll.actions) + + logger.add_scalar('gen/mean_episode_length', (~generator_data.ll.dones).sum() / generator_data.ll.states.shape[0], epoch) + logger.add_scalar('gen/mean_reward_per_episode', generator_data.hl.rewards[~generator_data.hl.dones].sum() / generator_data.hl.states.shape[0], epoch) + logger.add_scalar('gen/unsafe_probability_mass', policy.unsafe_probability_mass(policy(generator_data.hl.states[~generator_data.hl.dones], generator_data.hl.safe_actions[~generator_data.hl.dones])).mean(), epoch) + + discriminator, loss = train_discriminator(expert_data, generator_data.ll, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c) + if wasserstein: + generator_data.ll.rewards = discriminator(generator_data.ll.states, generator_data.ll.actions) + else: + generator_data.ll.rewards = -F.logsigmoid(discriminator(generator_data.ll.states, generator_data.ll.actions)) + logger.add_scalar('disc/final_loss', loss, epoch) + logger.add_scalar('disc/mean_reward_per_episode', generator_data.ll.rewards[~generator_data.ll.dones].sum() / generator_data.ll.states.shape[0], epoch) + + #assert generator_data.ll.rewards.shape == generator_data.ll.dones.shape + generator_data.hl.rewards = torch.where(~generator_data.ll.dones, generator_data.ll.rewards, torch.tensor(0.)).sum(-1) + + value, policy = ppo_step(value, policy, generator_data.hl.states, generator_data.hl.safe_actions, generator_data.hl.actions, generator_data.hl.rewards, generator_data.hl.dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm) + expert_data = roll_buffer(expert_data, shifts=-3, dims=0) + + if callback is not None: + callback(epoch, value, policy) + + return value, policy + +def rollout(env_fn, policy, n_episodes, max_steps_per_episode): + env = env_fn(0) + + states = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.observation_space['observation'].shape) + safe_actions = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.observation_space['safe_actions'].shape) + actions = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.action_space.shape) + rewards = torch.zeros(n_episodes, max_steps_per_episode + 1) + dones = torch.ones(n_episodes, max_steps_per_episode + 1, dtype=bool) + + ll_states = torch.zeros(n_episodes, max_steps_per_episode, env.max_plan_length + 1, *env.observation_space['observation'].shape) + ll_actions = torch.zeros(n_episodes, max_steps_per_episode, env.max_plan_length + 1, *env.ll_action_space.shape) + ll_rewards = torch.zeros(n_episodes, max_steps_per_episode, env.max_plan_length + 1) + ll_dones = torch.ones(n_episodes, max_steps_per_episode, env.max_plan_length + 1, dtype=bool) + + env = VecEnv(list(map(lambda i: (lambda: env_fn(i)), range(n_episodes)))) + + obs = env.reset() + states[:, 0] = torch.tensor(obs['observation']).clone().detach() + safe_actions[:, 0] = torch.tensor(obs['safe_actions']).clone().detach() + dones[:, 0] = False + + for s in tqdm(range(max_steps_per_episode), 'Rollout'): + actions[:, s] = policy.sample(policy(states[:, s], safe_actions[:, s])).clone().detach() + + clipped_actions = actions[:, s] + if isinstance(env.action_space, gym.spaces.Box): + clipped_actions = torch.clamp(clipped_actions, torch.from_numpy(env.action_space.low), torch.from_numpy(env.action_space.high)) + + o, r, d, info = env.step(clipped_actions) + states[:, s + 1] = torch.tensor(o['observation']).clone().detach() + safe_actions[:, s + 1] = torch.tensor(o['safe_actions']).clone().detach() + rewards[:, s] = torch.tensor(r).clone().detach() + dones[:, s + 1] = torch.tensor(d).clone().detach() + + ll_states[:, s] = torch.from_numpy(np.stack([i['ll']['observations'] for i in info])).clone().detach() + ll_actions[:, s] = torch.from_numpy(np.stack([i['ll']['actions'] for i in info])).clone().detach() + ll_rewards[:, s] = torch.from_numpy(np.stack([i['ll']['rewards'] for i in info])).clone().detach() + ll_dones[:, s] = torch.from_numpy(np.stack([i['ll']['plan_done'] for i in info])).clone().detach() + + dones = dones.cumsum(1) > 0 + + states = states[:, :max_steps_per_episode] + safe_actions = safe_actions[:, :max_steps_per_episode] + actions = actions[:, :max_steps_per_episode] + rewards = rewards[:, :max_steps_per_episode] + dones = dones[:, :max_steps_per_episode] + + return (states, safe_actions, actions, rewards, dones), (ll_states, ll_actions, ll_rewards, ll_dones) + +class OptionsEnv(gym.Wrapper): + + def __init__(self, env, options): + super().__init__(env) + self.ll_action_space = env.action_space + self.options = options + self.action_space = gym.spaces.Discrete(len(options)) + self.max_plan_length = max(t for _, t in options) + + def plan(self, option): + target_v, t = option + current_v = self.env._env.state[self.env._agent, 1].item() + dt = self.env._env._dt + a = (target_v - current_v) / (t * dt) + a = self.env._normalize(a) + a = a * np.ones((t,)) + a += 0.01 * np.random.randn(*a.shape) + a = np.clip(a, self.ll_action_space.low, self.ll_action_space.high) + return a + + def execute_plan(self, obs, option, render_mode=None): + observations = np.zeros((self.max_plan_length + 1, *self.env.observation_space.shape)) + actions = np.zeros((self.max_plan_length + 1, *self.ll_action_space.shape)) + rewards = np.zeros((self.max_plan_length + 1,)) + env_done = np.ones((self.max_plan_length + 1,), dtype=bool) + plan_done = np.ones((self.max_plan_length + 1,), dtype=bool) + infos = [] + + plan = self.plan(option) + observations[0] = obs + env_done[0] = False + for k, u in enumerate(plan): + plan_done[k] = False + o, r, d, i = self.env.step(u) + actions[k] = u + rewards[k] = r + env_done[k+1] = d + infos.append(i) + observations[k+1] = o + + if render_mode is not None: + self.env.render(render_mode) + + if d: + break + + n_steps = k + 1 + return observations, actions, rewards, env_done, plan_done, infos, n_steps + + def step(self, action, render_mode=None): + a = int(action) + assert a == action + ll_obs, ll_actions, ll_rewards, ll_env_done, ll_plan_done, ll_infos, ll_steps = self.execute_plan(self.last_obs, self.options[a], render_mode) + hl_obs = ll_obs[ll_steps] + hl_reward = (ll_rewards * ~ll_plan_done).sum().item() + hl_done = ll_env_done[ll_steps].item() + hl_infos = { + 'll': { + 'observations': ll_obs, + 'actions': ll_actions, + 'rewards': ll_rewards, + 'env_done': ll_env_done, + 'plan_done': ll_plan_done, + 'infos': ll_infos, + 'steps': ll_steps, + } + } + self.last_obs = hl_obs + return hl_obs, hl_reward, hl_done, hl_infos + + def reset(self, *args, **kwargs): + self.last_obs = super().reset(*args, **kwargs) + return self.last_obs + + +class SafeOptionsEnv(OptionsEnv): + + def __init__(self, env, options, safe_actions_collision_method=None, abort_unsafe_collision_method=None): + super().__init__(env, options) + self.safe_actions_collision_method = safe_actions_collision_method + self.abort_unsafe_collision_method = abort_unsafe_collision_method + self.observation_space = gym.spaces.Dict({ + 'observation': self.observation_space, + 'safe_actions': gym.spaces.Box(low=0., high=1., shape=(self.action_space.n,)), + }) + + def safe_actions(self): + if self.safe_actions_collision_method is None: + return np.ones(len(self.options), dtype=bool) + + plans = [self.plan(o) for o in self.options] + plans = np.stack(plans) + safe = feasible(self.env, plans, method=self.safe_actions_collision_method) + if not safe.any(): + # action 0 is considered safe fallback + safe[0] = True + + return safe + + def reset(self, *args, **kwargs): + obs = super().reset(*args, **kwargs) + obs = { + 'observation': obs, + 'safe_actions': self.safe_actions(), + } + return obs + + def step(self, action, render_mode=None): + obs, reward, done, info = super().step(action, render_mode) + obs = { + 'observation': obs, + 'safe_actions': self.safe_actions(), + } + return obs, reward, done, info + + def execute_plan(self, obs, option, render_mode=None): + observations = np.zeros((self.max_plan_length + 1, *self.env.observation_space.shape)) + actions = np.zeros((self.max_plan_length + 1, *self.ll_action_space.shape)) + rewards = np.zeros((self.max_plan_length + 1,)) + env_done = np.ones((self.max_plan_length + 1,), dtype=bool) + plan_done = np.ones((self.max_plan_length + 1,), dtype=bool) + infos = [] + + plan = self.plan(option) + observations[0] = obs + env_done[0] = False + for k, u in enumerate(plan): + plan_done[k] = False + o, r, d, i = self.env.step(u) + actions[k] = u + rewards[k] = r + env_done[k+1] = d + infos.append(i) + observations[k+1] = o + + if render_mode is not None: + self.env.render(render_mode) + + if d: + break + + if self.abort_unsafe_collision_method is not None and \ + not feasible(self.env, plan[k:], method=self.abort_unsafe_collision_method): + break + + n_steps = k + 1 + return observations, actions, rewards, env_done, plan_done, infos, n_steps diff --git a/src/safe_options/policy.py b/src/safe_options/policy.py new file mode 100644 index 0000000..4e0e9d0 --- /dev/null +++ b/src/safe_options/policy.py @@ -0,0 +1,35 @@ +import torch +import torch.nn as nn +from torch.distributions import Categorical +from torch.distributions.kl import kl_divergence +from core.policy import SetDiscretePolicy + +class SetMaskedDiscretePolicy(SetDiscretePolicy): + + def forward(self, observation, safe_actions): + return torch.cat((super().forward(observation), safe_actions), -1) + + def torch_dist(self, dist): + logits = dist[..., :self.action_dim] + z = dist[..., self.action_dim:] + a = super().torch_dist(logits).probs + return Categorical(probs=a*z) + + def unsafe_probability_mass(self, dist): + logits = dist[..., :self.action_dim] + z = dist[..., self.action_dim:] + a = super().torch_dist(logits).probs + return (a * (1 - z)).sum(-1) + + # def torch_dist_nomask(self, dist): + # print('no mask logprob') + # logits = dist[..., :self.action_dim] + # return super().torch_dist(logits) + + # def log_prob(self, dist, actions): + # return self.torch_dist_nomask(dist).log_prob(actions) + + # def kl_divergence(self, dist1, dist2): + # d1 = self.torch_dist_nomask(dist1) + # d2 = self.torch_dist_nomask(dist2) + # return kl_divergence(d1, d2) diff --git a/src/safe_options/policy_gradient.py b/src/safe_options/policy_gradient.py new file mode 100644 index 0000000..6e968e4 --- /dev/null +++ b/src/safe_options/policy_gradient.py @@ -0,0 +1,107 @@ +import torch +from core.value_estimation import gae +from core.optimization import conjugate_gradient, line_search + +def trpo_step(value, policy, states, safe_actions, actions, rewards, dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1): + + states = states.detach() + actions = actions.detach() + rewards = rewards.detach() + dones = dones.detach() + + advantages, returns, valid = gae(states, rewards, value(states), dones, gamma, gae_lambda) + advantages = advantages.detach() + returns = returns.detach() + + # update value function + + for _ in range(v_iters): + v_opt.zero_grad() + value_loss = (value(states) - returns).pow(2)[valid].mean() + value_loss.backward() + v_opt.step() + + # compute policy gradient + + plogprob = policy.log_prob(policy(states, safe_actions), actions) + surrogate_advantage = (plogprob * advantages)[valid].sum() / states.shape[0] + g = torch.cat(torch.autograd.grad(surrogate_advantage, policy.flat_param)).detach() + + def Hx(x): + kl = policy.kl_divergence(policy(states, safe_actions), policy(states, safe_actions).detach())[valid].mean() + dKL = torch.cat(torch.autograd.grad(kl, policy.flat_param, create_graph=True)) + H_x = torch.cat(torch.autograd.grad(dKL.T @ x, policy.flat_param)).detach() + return H_x + cg_damping * x + + x = conjugate_gradient(Hx, g, cg_iters) + npg = torch.sqrt(2 * delta / (x.T @ Hx(x))) * x + + # perform line search + + def L(theta): + rplogprob = policy.log_prob(policy(states, safe_actions, flat_param=theta), actions) + return ((rplogprob - plogprob.detach()).exp() * advantages)[valid].sum() / advantages.shape[0] + + condition = lambda theta: policy.kl_divergence(policy(states, safe_actions, flat_param=theta), policy(states, safe_actions))[valid].mean() < delta + + x0 = policy.flat_param + g0 = torch.cat(torch.autograd.grad(L(x0), x0)) + theta = line_search(L, x0, npg, g0, backtrack_coeff, condition, max_steps=backtrack_iters) + + # update policy parameters + + with torch.no_grad(): + policy.flat_param.copy_(theta) + + return value, policy + +def ppo_step(value, policy, states, safe_actions, actions, rewards, dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm): + + states = states.detach() + actions = actions.detach() + rewards = rewards.detach() + dones = dones.detach() + + advantages, returns, valid = gae(states, rewards, value(states), dones, gamma, gae_lambda) + advantages = advantages.detach() + returns = returns.detach() + + # update value function + + for _ in range(v_iters): + v_opt.zero_grad() + value_loss = (value(states) - returns).pow(2)[valid].mean() + value_loss.backward() + v_opt.step() + + # update policy + + old_dist = policy(states, safe_actions).detach() + old_logprob = policy.log_prob(old_dist, actions).detach() + + def g(advantages, clip_ratio): + return torch.where(advantages >= 0, (1 + clip_ratio) * advantages, (1 - clip_ratio) * advantages) + + def L(states, actions, advantages, clip_ratio): + return torch.minimum( + (policy.log_prob(policy(states, safe_actions), actions) - old_logprob).exp() * advantages, + g(advantages, clip_ratio) + )[valid].mean() + + for _ in range(pi_iters): + pi_opt.zero_grad() + ppo_loss = -L(states, actions, advantages, clip_ratio) + ppo_loss.backward() + + if max_grad_norm: + torch.nn.utils.clip_grad_norm(policy.parameters(), max_grad_norm) + + pi_opt.step() + + kl = policy.kl_divergence(policy(states, safe_actions), old_dist)[valid].mean() + if target_kl and kl > target_kl: + break + + print('KL', kl.item()) + + return value, policy diff --git a/src/safe_options/test_options.py b/src/safe_options/test_options.py new file mode 100644 index 0000000..269eb43 --- /dev/null +++ b/src/safe_options/test_options.py @@ -0,0 +1,54 @@ +from intersim.envs import IntersimpleLidarFlat +from options import OptionsEnv +import gym +import numpy as np + +def test_obs_shape(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + assert env.reset().shape == (36,) + +def test_act_space(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + assert env.action_space == gym.spaces.Discrete(3) + +def test_plan(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + env.reset() + plan = env.plan(options[0]) + assert np.allclose(plan, -13.998268127441406 * np.ones((5,))) + +def test_plan2(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + obs = env.reset() + states, actions, rewards, dones, plan_done, infos, n_steps = env.execute_plan(obs, options[0]) + assert states.shape == (6, 36) + assert rewards.shape == (6,) + assert dones.shape == (6,) + assert len(infos) == 5 + +def test_step(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + env.reset() + obs, reward, done, _ = env.step(0) + assert obs.shape == (36,) + assert reward == 5.0 + assert done == False + +def test_ll_step(): + options = [(0, 5), (5, 5), (10, 5)] + env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options) + env.reset() + _, _, _, info = env.step(0) + assert info['ll']['observations'].shape == (6, 36) + assert info['ll']['actions'].shape == (6, 1) + assert info['ll']['rewards'].shape == (6,) + assert info['ll']['env_done'].shape == (6,) + assert info['ll']['plan_done'].shape == (6,) + assert info['ll']['plan_done'][5] == True + assert info['ll']['steps'] == 5 + assert len(info['ll']['infos']) == 5 diff --git a/scratch/etienne/trpo/wrappers.py b/src/util/wrappers.py similarity index 100% rename from scratch/etienne/trpo/wrappers.py rename to src/util/wrappers.py From cd58ce289838e3eb47a8af27cf3fb54100f3832b Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Thu, 17 Feb 2022 22:43:41 +0100 Subject: [PATCH 06/10] Remove old code --- src/eval_main.py | 2 +- src/evaluation/evaluation.py | 2 +- src/gail2/envs.py | 111 ------------------------------ src/gail2/options.py | 128 ----------------------------------- src/gail2/test_options.py | 54 --------------- src/gail2/wrappers.py | 74 -------------------- src/options/envs.py | 2 +- 7 files changed, 3 insertions(+), 370 deletions(-) delete mode 100644 src/gail2/envs.py delete mode 100644 src/gail2/options.py delete mode 100644 src/gail2/test_options.py delete mode 100644 src/gail2/wrappers.py diff --git a/src/eval_main.py b/src/eval_main.py index b77bf91..6b460b3 100644 --- a/src/eval_main.py +++ b/src/eval_main.py @@ -10,7 +10,7 @@ import src.gail.options as options_envs from src.evaluation.metrics import divergence, visualize_distribution from src.core.policy import SetPolicy, SetDiscretePolicy from src.core.reparam_module import ReparamPolicy -from src.gail2 import envs as options_envs2 +from src.options import envs as options_envs2 from typing import Optional, List, Dict, Tuple import torch diff --git a/src/evaluation/evaluation.py b/src/evaluation/evaluation.py index 20a6db9..27817ff 100644 --- a/src/evaluation/evaluation.py +++ b/src/evaluation/evaluation.py @@ -6,7 +6,7 @@ from typing import Callable, Dict, Optional import os import pickle from tqdm import tqdm -from src.gail2.envs import OptionsEnv +from src.options.envs import OptionsEnv class IntersimpleEvaluation: """ diff --git a/src/gail2/envs.py b/src/gail2/envs.py deleted file mode 100644 index 6f6bdd4..0000000 --- a/src/gail2/envs.py +++ /dev/null @@ -1,111 +0,0 @@ -import gym -import numpy as np -from src.gail2.wrappers import Wrapper, Setobs, TransformObservation -from intersim.envs import IntersimpleLidarFlatIncrementingAgent - -obs_min = np.array([ - [-1000, -1000, 0, -np.pi, -1e-1, 0.], - [0, -np.pi, -20, -20, -np.pi, -1e-1], - [0, -np.pi, -20, -20, -np.pi, -1e-1], - [0, -np.pi, -20, -20, -np.pi, -1e-1], - [0, -np.pi, -20, -20, -np.pi, -1e-1], - [0, -np.pi, -20, -20, -np.pi, -1e-1], -]).reshape(-1) - -obs_max = np.array([ - [1000, 1000, 20, np.pi, 1e-1, 0.], - [50, np.pi, 20, 20, np.pi, 1e-1], - [50, np.pi, 20, 20, np.pi, 1e-1], - [50, np.pi, 20, 20, np.pi, 1e-1], - [50, np.pi, 20, 20, np.pi, 1e-1], - [50, np.pi, 20, 20, np.pi, 1e-1], -]).reshape(-1) - -def NormalizedOptionsEvalEnv(**kwargs): - return OptionsEnv(Setobs( - TransformObservation(IntersimpleLidarFlatIncrementingAgent( - n_rays=5, - **kwargs, - ), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) - ), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5)]) - -def NormalizedContinuousEvalEnv(**kwargs): - return Setobs( - TransformObservation(IntersimpleLidarFlatIncrementingAgent( - n_rays=5, - **kwargs, - ), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) - ) - -class OptionsEnv(Wrapper): - - def __init__(self, env, options): - super().__init__(env) - self.ll_action_space = env.action_space - self.options = options - self.action_space = gym.spaces.Discrete(len(options)) - self.max_plan_length = max(t for _, t in options) - - def plan(self, option): - target_v, t = option - current_v = self.env._env.state[self.env._agent, 1].item() - dt = self.env._env._dt - a = (target_v - current_v) / (t * dt) - a = self.env._normalize(a) - a = a * np.ones((t,)) - a += 0.01 * np.random.randn(*a.shape) - a = np.clip(a, self.ll_action_space.low, self.ll_action_space.high) - return a - - def execute_plan(self, obs, option, render_mode=None): - observations = np.zeros((self.max_plan_length + 1, *self.env.observation_space.shape)) - actions = np.zeros((self.max_plan_length + 1, *self.ll_action_space.shape)) - rewards = np.zeros((self.max_plan_length + 1,)) - env_done = np.ones((self.max_plan_length + 1,), dtype=bool) - plan_done = np.ones((self.max_plan_length + 1,), dtype=bool) - infos = [] - - observations[0] = obs - env_done[0] = False - for k, u in enumerate(self.plan(option)): - plan_done[k] = False - o, r, d, i = super().step(u) - actions[k] = u - rewards[k] = r - env_done[k] = d - infos.append(i) - observations[k+1] = o - - if render_mode is not None: - self.env.render(render_mode) - - if d: - break - - n_steps = k + 1 - return observations, actions, rewards, env_done, plan_done, infos, n_steps - - def step(self, action, render_mode=None): - a = int(action) - assert a == action - ll_obs, ll_actions, ll_rewards, ll_env_done, ll_plan_done, ll_infos, ll_steps = self.execute_plan(self.last_obs, self.options[a], render_mode) - hl_obs = ll_obs[ll_steps] - hl_reward = (ll_rewards * ~ll_plan_done).sum().item() - hl_done = ll_env_done[ll_steps-1].item() - hl_infos = { - 'll': { - 'observations': ll_obs, - 'actions': ll_actions, - 'rewards': ll_rewards, - 'env_done': ll_env_done, - 'plan_done': ll_plan_done, - 'infos': ll_infos, - 'steps': ll_steps, - } - } - self.last_obs = hl_obs - return hl_obs, hl_reward, hl_done, hl_infos - - def reset(self, *args, **kwargs): - self.last_obs = super().reset(*args, **kwargs) - return self.last_obs diff --git a/src/gail2/options.py b/src/gail2/options.py deleted file mode 100644 index e2d7263..0000000 --- a/src/gail2/options.py +++ /dev/null @@ -1,128 +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) diff --git a/src/gail2/test_options.py b/src/gail2/test_options.py deleted file mode 100644 index 269eb43..0000000 --- a/src/gail2/test_options.py +++ /dev/null @@ -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 diff --git a/src/gail2/wrappers.py b/src/gail2/wrappers.py deleted file mode 100644 index 3916e75..0000000 --- a/src/gail2/wrappers.py +++ /dev/null @@ -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 diff --git a/src/options/envs.py b/src/options/envs.py index 6f6bdd4..6c58045 100644 --- a/src/options/envs.py +++ b/src/options/envs.py @@ -1,6 +1,6 @@ import gym import numpy as np -from src.gail2.wrappers import Wrapper, Setobs, TransformObservation +from src.util.wrappers import Wrapper, Setobs, TransformObservation from intersim.envs import IntersimpleLidarFlatIncrementingAgent obs_min = np.array([ From 9de6bfe9a38b0d6d2b6c5d7304af23f30aec171e Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Thu, 17 Feb 2022 22:58:00 +0100 Subject: [PATCH 07/10] Add GAIL --- .../gail-intersimple-setobs2-03-02-22.pt | Bin 0 -> 14187 bytes evaluate_models.sh | 3 +++ 2 files changed, 3 insertions(+) create mode 100644 checkpoints/gail-intersimple-setobs2-03-02-22.pt diff --git a/checkpoints/gail-intersimple-setobs2-03-02-22.pt b/checkpoints/gail-intersimple-setobs2-03-02-22.pt new file mode 100644 index 0000000000000000000000000000000000000000..1a9454705954c0c8e3ecc437b50eebc4ecf719b7 GIT binary patch literal 14187 zcma*O30O|=w?EuGY8Ik7rCEcqXRo`2NOOrwO3{R*$viaCAZbpLXi_Q>qMp6(NTFzu zDVdcbnKFc^f8X;v?|IMfyyt&i?|ohS+I!z?t-bf!>v{GX?)7oB;TPcH5fS0}{|sp! zSstIT^?^a#{AO?PiSRMpx@ohOr6iBXzl|uq_3J}6Z}wXs5fri|T+AVCgI}26h9yDk zBg8^^beuzZ-NYuY?7!J3!h5Svm``vh-};D zNSOCF6R}W!ZEZK<390Z2PM-k3P=SAmh*@n8@riKx*Q8jepd0T5k5Q&_J%6_ZxrLu>3c(G{BIPKP{qBWO8*qI zH&pFED3Rg5hK@GEBBCo(B623)(LbBx|JBTrJO=+}=GjJ8mOB5#$p2>Xe?Q4RbDK2k zmcX^$g`_6?6+F0k7$haukwXCnB#?Ur%`S@5TXQ|hJFg59|Lh&Ukl)8BIxA8COZ>R; zr~-K9+QWtAXYip8a)7OdLy32>Y-BF+2nnTcwjihGOgb)bY+|#@6EQLFIDNN8 z1f@bM;E3y6W@M@{T{)Ocub2AMf&2`5lGDqc>^np{XU3AJP78>vauR(ydYTC3AEuMG zUc&9^u{5J^8kzN2j(ld0QaF79e4;6PTyz%9+Z;;;Qqrj~IY4i(omvJL8% z$H4gBcY1sNT3Y+T9wTHpH1=aD)d_q@^=5|A*7Y*<+~n=F#{UwG9a%*;TV234L1(dQ zmOh<1O{wAOm&Hs^Ry%$7aR9V@3@N{z0o2D8kg-#q*pso72<%@7#bk<(duj zueV_IlvdFQ7f0H)TaSLVwxt^$XqazZtj(gp8v5gn2AOJJfe%h~L5H^zoiZ=HVQ;5T z!+F<2{92uW;)fTZ-QfTt{1vhFz6E>eO*^(EiGqIS9eO!An8+Ns4`LT}=(~>jI6JM8 zPCnR6E$7D44a!=uPlk`q@+-wz#YHr8_X?`AZV|ckp_xNWkI-H34-@qm3pkw;NcYJ~ zQhPTCh>^_#_iw5+fD=k~sEN?wKa0q@hhikRi;FT*40W=YN?o@EfD5OBYOmnKqBoPM zveYCp&~*~DP8iTi_g`3Ao($rw1gxo?Lpzs9QQjF%(9Qn`bE0a=kL6PEy)}X`bN8}> zyptfnrXEHb$}nKX3NlM`Ia#^)6p41zA?mZEXxB3-a#UwA+{}|DK8_L)tc9?)J(>LQ zs3E2665yD70@}6jp~w8t{@O{GY$#dqEy@;lUyF9XP}` zn&y(KWH$<3y5#$Y654w$jy6u4N=~~Srq#s(lpSq?i+2i0!oq%5_Z<&W&7M#ER2Iky zvyzUk3MMwf6;vgw5kpvA+Ii_TYp`e!YKJ_6r^Om%mwyPj4^~6fvJU#T%$20L9Hv%x zSaz^B8`~Uu&_XVkn2ssY71_s0vWJ$LoR=80H7EzbM4$RZ+(Q*AM3QuVvv*%EV~j6Y z;Q1^?dRY1bboeeOrV*dul}i(LY>^^Ug^rLE%OGNsZv&3|lc0KcAG5DkhMdsOz=+dX z=vSmnBqlph-Z{#2*UU^%c)!qmrfCurty@OyY(D|WmXX(evmwQH82v1G>9C?6M(@q0 zSM$m!&n#oMJjWmYep4lqG1xi23*> z3NaBRVP^rYQc@-}P7Py);B5SMQJBtE+e%d(KjMq~l~iS{2S=7(#pfa`nb@BfsGZP$ zDlcbAyOJHq7oLr9*LyFH@x8(YQO0E9_xa?nUl|r{o=qYK!eQh>5&0=yNRzHaQSOZe z6s`Z_cxyGCZ!sCWSWVhFB8l45Q$Q@xmi}CMj_|!6CnaOMFw0FJE5GW{-F*Tyr}Pll zVr?w!VJykPlC!k4S(Z-D+CfhQ%>>}#!<$#X0=p~|jTU^}(zv8K8&J+Y}-xFLT*G7bGG2v+fHNp%)bAGVZz{>+AUFRg(ziz=JUQO{ z8QxcJgvXxQc#Y%$!T+I!O?J2}4$P9w9St%hKv)$wcp349$(VLE9tq zh^Ffb@^1VHyKzY>+V;&vbYDQPNGmeyM)IlYmk#uLc!Jd|D`c*QmCyqr0d%pO4~-2v zjVl$qm?hbViBz#Rd8v~Dd))Pj^N$eh&ncmMXKllP;BdNY`41X>u7~b7sfUhhfL&rC z_)@oq#!5G0&E!Hld`}+@?sietp{E$vZ%uRdd6DeYdZNi|Ku>K?C!WD~KsIq0bH$dC z?G>}Avz0KFHKDX+RtBm2^a!-Fld0Uz#dL{WDfmTBCM6!Hh>BJTIh#43?#<}L+~o|7 zZ%ZbZJk0U>by+fB=Q-poKF=JuF930GenZpUE_RxYIL+}b#Ol@aNtE0%>OWbUtSi-{ zi6Rp=xw4Pet2U6i&SI3`WhwDu7Z4o2K(`82(sOayM3?5Wg6btGn>OLg-OGuS>o;cY z2+-QND40ERfwkzbqldKup;oL4wF|9jRQx_tII|2F4up}Fj~hv9gF6hx_T#)G>7-BP zD$Le7M)P&fQ_sD{bk~MlRy;Z%n@;$_ZiiS@eEx!?8LC0A^ztz=`30!&RDza4FHo%P z5}H&lg1!qb*na5%d!UHIVHrKzp7Iwg0`|bFbJIxP143^Y4{$hVufl7MFO2fgF~+nz zk9F2Ri2K&7aj$Mj!nrGwk@NOC6CS>TjO`u;h4bQ=_n;mp86IIO3Iy@}bx|_zU5(@G zH0WcIOV}GiVbFb;*|$?3Y^vUaiQ5%8)oRbVFZ>F-Ry_y5IVKpj!VmH%bEw#*9e8Kf zP2>?j1Q#`U$?FUGOn?DplBc$yskHk~QR<2*K954?%q?Kgjs<61DA}_`;V$dAB5{n+~$C zIQy|@z6!4Tp+pKl#lV$&GStsO8-G|=fYIG%M!GPC@}`n82@kJC@HgdKSty8>gnjaO%+D0ufVcH2SBkb z1vTeOl8W&Hkl9tjglTkwu0th!lV6J!Rxi27K1D;{eqN^I!6Jy1A*jb!fh|KtFvH{` z%o|X`p*5q3>`l%pc}sBgjfa_sZeYiP3}#--Ihd8biXEKw9%3F9L6f2aEB{KDE(Lii zvY-kK@^2tFq?Yv%I0cCYI^(8iu#m`1u$VW}dhu!9(O^EmDa{my9AUBU}V^VsO^ z19)1p2lXfXe0j((6iV6)Q!pDn!fQEu3i+vbv;ntshY2(D_6?Z5sE7SE`~#wdvvJ^t z9vM(P4|A$o@p_Ox^yLZ=ALcwOsabqJYj}YB(>oKmL@@^uwo=Rh5 zbI^A~3{D@+z~xr}f(;+Q^a&rm49;*OGoHe92^r33HzzhI^f^Yk&!;iJ)}r^pH_Wv& z-y!y|HCzdc2l2e~Ai1*+nc`%)=Oo9``CH69EPP+@t~Qf8#wg;FL~(k3#xZ;vp~OCX z{E69~)rH0@<#4r_2x&R!iM<^{)b3C=s@^++A1nOP?r`o#@4{(oZlh4FAk~(-0k+~&B;(^;ddfnD zT^-+#J;5u{Fhz&(SMgEflchMwO2hO$`eeagRig2=9K*NBkhf$Zr|6Y9ER3|r<@3+L z=|5b!w=)!XJToJ5vz2jVU?okM0r)bD12x^6q_S}m25fF&n6T?u^mZaXI!RfzNeh_J z(97`T_G{Q`ngEfe*P+w71{Ag*MGIST&~%Z&?|-y$N1z8S3;lylYkBD-t$LiY&j)A4 zs?n^!&X|>`j_0Rzaa@niBy3$P^VvdzMjV+zu5mOOjXlTeZ|Z4)f1Wbz6F}q-79zh~ z0!aG3Jn$}Opu3Wn99-UmN^e3~5xMQS>%wMeovu#5-CqF16LG`3-6mL9z6=@!+raT^ z+x!jbq4;F)C3caJ0L|QQ3`eu3QIo%(RP}cn3hlgzMn()uiHXrV>pbvE@uXv?WHG?^ zI4$P;%C_u?;bNmQ>6m_$VQc1MieVt}4eAExFRS3EU>bC*E5S4Q%kZW~lFV!V06f7F z)c?;qyvVm36o03{r>RBYy(o%~IVn?zkNUJwNe@GwZ6~tK3KFvC0LhHu!BvqgmXv7H zg^f>9a%YIc49^@4X{O?AcE9T-T&BWd zr{D9%I_W5k>fc2l@(L1VFD~ArBaHHwU0^A6o%_y9pIoz*Ws>9fK$+rgc0tx0O84$1 zX6@}xd#{=ud^#QMGgxNvgS8-6*bT|+E|9dXpWsyS5y*Sd1I7<-LQ`-r zC{sDkukjcxi;^IGt+U~%(*X9@PQ~NpQLw>&h$C)r8O);pqT|41+H&&_h<}wL+Dg^z zoBcIV892_gM6N_7(K}GXy$F4WpE9D;q}d9KpZI=W0=Sg!1-F~$;0KQonKC;aZY$)X zN780U5WbD2UUjfcI~y-MK7yFEg}R5zz;P2PTqtoApAU~?&XTXd z``ZHMyncZzvO4f)S`ORKYsSn{KgO1Hijh|y@44Gvs6h0ujc}t)k<4p+!2Z6(N3EVp zQl0fVxZ-F%S4-+Qqs=Wa4wf9FE!k zi>%@O1TcD&0>)p9(dou0m|gz?rw^u}b_#{dk-_U$6j*c&j03a ziC>R$28ztgHFwZ*&p`;Qp9kzaa}r+}SYDu5m<-w)3DO zpNzXi;~Bau9p8@?;`JS0P}#2$#kc111=@TNQa`ouBw^ zMZ)_MeJJDK0hPDULg8I^4AoL1yb8f^!L0^oRo#NK*HX~*PZcb&xr24{GRsTiS&b@E4VH6Fb2MKKH`T*m&gn;h4Mg;LYOef94N@3aD*hy@TDvWUQA+cvJ0^}%@ETCWPuy72LJdh1eH4v7!>%-&J8*XzA4u+ z^5j+4sHKx@VE=*{g>HFK z_H%>+iEC9MS6=7goc)>1%v)SWDKMLDntBp{uUgD{rPYF$XMkC#QxG?bGsb*vG6mgR z^3*yl8MjRQx3!Kl_~0l*y!Oa5_ZVriLn@V)wtQfAD+o~E(k-CbQ;u7;yE)<3vgC4W z6Jw&Z2N@e)sS#|tYE-J61`6ZvpjC1n-UQn21)7|r+RqiN`D&~gZd zs;(T^Xst<9>(+v;>V^7Vy}KN1e=X2VmhF;8`UoR(FP-(^<)<4?h|s$26Di>p1zc#bA0%V$qR8Xd_+o_!@jaB!PLFWIx(}`} zSTdI_wp$Nl)9=F@;SBUwSpc`*CgQ4+_t>*1aN=9mhEMigKyFDbyiJ@!T8rBeQw!&TdHaat-i zZ#~7?G37gMI2q1ZTD7u`q3>~Iwh3wUs4z3>TSjh(Yf;_YAMm@X4Bp18k(@b`xGIwN zbXVgoZ0zm=nKgl+prS}Vmsg|L$3!Tr)}(JCztA>FhAQs!hFufiN6DBju21QC zh_)+(WLq_MpM?oLd(TTh{#gY<_vA>(X-DvwJBipPYqJ|YQlPg$62cxIhVa<~AbwPo zt}Oco-34c%d-@|7Khn(TrDlR|%`;|Pq8^KH3uRq)ZDmNL|J~Vfhfsca<$%VuyzwETpqj8VPhK-7Y$h*drno04{ZJgWO6Q2di3kiy&N08YO`~YQM41u~)FDnq$h)K;_ zbgE)6(B6L7pjChp*BDmq`vchV`#v}MSr6J7 zqtR#BK6HtZ^t5NsEkdFrM=BZ#jx*YiWoi4@16chUAVKtdok_V0wVvw)8pDOmf-*&-^+}M3cV@%=p;^)T-?Ymc=#CCJh#Ic@8zky4`GR70E}(Iu4|?w1LhqC(uqiwV{;bKy{pQm#R8)nyL}h|y&Ud!U zFc4OYG{Ti{UvP1(JjKrfFz@6Ha?CRUF6?PHv)&^{guLD1$mJL8@Y7~U_C1D|=O4tR ziUMqHuV4)vMM(d%!?5pV8Jo^E!xB#kR%opyV?MBgX?FjOZl-79$E}YrH*P{*^E1S4r|`j#^_cJn7?00!@N&_5_$~2*y}~5oH3>1?ugp)!B%QJ5E(;GX z&4EcX#PPk&PAGnm&V1q=Wf7-Pj~)4__0tJ2x8J_G1<8RF0sw5h~!;p_dExaex{@Bo%6VS(g34!{36`CqrjBzP^C2qb=djD z3Y?pBapTH$IB53-&x}q5x^@@}l?y=Jy9u6qz<0sAa4*7@xnLKDt!b0t z;HT#(U|fiL*FQ1mj>?e9Ti&9o?ifm^HQ;s5PiPq6Cl}7&#s>?gk*4+)OvU%@5W4j` zYIv*3a`Xa2@;tp~pd!Vq~l<_uAgY{{0*d`W$ zE80?UkJ445%Lj z8D|&hvoD9OF^bfn?gaaQw}>5Xh$W@7M_D0{XxK5M82PsNqh@XzWJ$+DwMjF$irvR; z`Kp+jBu91x&7)fc)Ih1li0DOjA}?PJc(gZTqtiGmlcgTFAO}3ug#-v+KtQO&yuY5g1cg?rZ?*)#9v`R|Ld7UxtEfpMkwcKk%3~ z!7Kf691HBm)xo^vy38oNGH0Pum7h+s(}sFIP3ZBp!l8CucGHO_NR3x!4_>&Bs?-Bd zZk$YZ+8ux|>hh%Md?QXO2<61<_Tt&P1lHkL7V-~m#&5~DSetuyFxTlDEZ3S!Il=lE zax4UPcPL|H-e%@k^FlVkxB_3Bmaw{m#aJabhx>Q(MU1iFK;O)2>}i_H%(|OMVvU<| zyFfh!6S9^L83l`BlC zio$ei^ynEb&$t3vdlaeQ{GI3{GoK7dG{BsdPS|`Y6=N4h!|~v!C@i_o>_lY;KE9gC zI$sc^vtPR6wvaBgU)sUyOg@6=cuU~zhzhwi)0w2zjlzECN3bOR3_X$-N~G2FV64R! zrz%V%+5vpDZBYcAu5=;aU)51(#+w|Ol1G02JxNw)jpNS6lZZE;JmmP!!X5oW#3s9j zsQ&%{p~2J0hx|q?T$_iPYyt^V4q=urY(c-?317K!fYC^gA_m)?>BB<_FeLdJETR{p z`{4xgGI}x5Yt*HtMcQbx|0?<9w3dYFacT6zG`O|&1KC#`3@$ghsF!hsmbxq=GSg1b zkg1j!q$dQU%GLCuC63c^k*)m!XxaH9nWFcx?x=M{v5a-nnt4HX40(uyKMTUGTQPa2`889l96Z2h{Jvt zI#RI%oPId4BPuT-t%jd&{W={srk6qLguA+YVKn~xFqd867kxC2J7o+QnJRGw-0=(z0;kd35Dd*2ZI4MnAm#zRWUokp`?;+D_ zCQOh26aqYZgp6?ns0Ncv*7oGG_2=%OL+lHrA6fXGK|tKRiD~#OLN^I~VDI0Y&`XVl zg5d<_?oSnFnYlJ*guMZe4+r7lfoJF?-w9s3jga?UKV(_6kAqiHjs3h{nJcQrl^S--NsKhi_Jo2xSe}($ph~7voYNBg>ASjq6rVY{R;1X zp8&u6FVX$@Ah&Zw2G2%#qJMA=ic8&xn(d3(wDp2irOF>{ipNm>vp$$ycjw|_Rpc)T z=Vnw6g8My5(*I4Doc=tQ#>bRHo_Z7*N~s#z@OJKS05;^W7Gay~3 z4L_;bgIAM2+&TS<$#vx=F)_K|c_Rs{=jDTP{XR%=6{dYVCGdHpF4&uoF?|QyV1|V_ zJ@sKXK3iUlN{kv!oRz~WpP8_EVl8tbb)tqaR01<>@}ceP7&KXW3z#!h$;vhS+kF^P0| zTf7E8y5=%{(ce(F?J=enp9N<5oRfU`4+#cfwKTyPTP zj!M#n#2JjzzM;@a_(WaP4Zme4V!-V!tWPZ;eekOSyI)(uN3%Cj^f-s%GveiL^2=kE z+I`0pPJi)-umssTQ3snQcns_;YnfXw^VlU)YPi+mG}aW1p>a(dXWfP_xPIXRTueB` zY385M)9WwH_v>a~D@u?Pd*4s=hO~p#+H0UAk_XveUjxT548n$ukw0x5r(8%ymBp`c zgS!M?*Q@}|;ateJvj;^@J2Y#Uin~1BA-8QlL>O#^wb^o%-;>2n^a_@b|6u|P8nJJ3 z0BdybFD`!52cOqSQ!DdgEcE0huJ0AtBT@$V?bA1WICTJAE?2Ud?PJV3^HWUgq{H=b z#%_%MNnO~RBuAwh+8GDaB{Zbl8SlUI#Ce|&Pt@;nFmbgVoJ(5`+1@qSx#cvhJ1$F2 zCKQ9=EMeH>s@8++W-N`uXzOR z^>XBq(n9v;HALPj5%|G5;Vs}7y|dnQ@xj=Bx^1oiPSBC@ipgR21*guK8hJIGa!kPb*z-j7<|kXBtg76 zc;{6zEM6}~3mbDMbX^B0JalRA)PD48ai(vV#gW}kU+X&(bwQ>t6)S&Q(<$=jNrcjG zCLzWQT7Jt=m2p@4>xm19-1v&ax3b~P+rxmjzcTCc%*lL}MVO+q0TlZCaKo@CJWzfM zKeiM>;S5JQOa2LVJ+oxiJu?R>V>Kd?R)(M2RRuxEUMmdx9FN}qax{CFvcoIk;X;FbRi z%XM!uYs&@Ts>dPR+&aKkH}(EciIZn&3F@Kp}ZnFCv zP<^Jx9?&nrD$lT5=^>7uE5YdO z4os~}q1{5UAn4-9*i>}0`SYK$(dMrhx0ZU8-;#v>`zCsU^xI*>^9ttkN_YG)^E|W8 z(Ug2%(r@;nM~Y)4qk}DvSApkqSDZd7L$jnMh@tft_Uge>RC{?2wYMC`4+#^!bcrcU z`m4?7{2723Tr`(q-yD4-t%DW(-$NjTr(!Q_yjxC z{TYNyt!Aw4@|YOEN8I%H{cQH->EJl86Le;5Wu|YLijAG&MEj=_eV18m7V!K!6SQ*$ zv+jLLJ&B9p8js5pdDa3W&zR!yvcidcaua5ymE(tC3mlu=jH*X&zd{9`^2bNzZ)~ekbS;1KvWcTiBX8i0kEd8Cy z{^Gv{S%u$$W8VhM*H+M-HwE$~zi`iNTSYzdr(l2OX<)^|G0fkW&52T{QHFbP;|?Xd z)2ACF(3I?4>W;7M-*8iu1EE^ek}RL2hS!_rNK!rzQ4-F8l~=WxRig`t`V2qnn;bxN zywyOXMFoHC5@ocW3t_` z&ouT+LoMi+$+L2PhUlxE1^Z?nME&+MZ2NN+q9^)N*F8{R%`Z)?A7;i7^=jGru=qt7 zoc9#jy<2gT_f|YJ5hs1UJQV{E`qw|#%w)`7zC`oh=L}h1hI^&#VQi5&36qdy&HL+d zRtOrzAPzR*cbq&%xXR9x~Hfmh)F@2DVHq#)7qzh^R^#_uE4UIQr!(X7^Px zKLjd~e}xT*hLu8bPcGAblb==3YlQH{OL1*|EE64(1kalmfJMO?5UQDn>5kQiyO*(Z zo3FsW-)&g@?G(EBO(T`6L2QeU9X6lXiD`=2++8JnR60Hg1!8!}ihH#df-5EawuH7tDd?G439R6g4I;*NpGnKuqzxkGoYyQSXrg z)r&pJKH=VG>mT|+*?fOAer<|>hR?!uc@gqf*BTdPxHH-5>TF=gI4cqP4jRfu;Y>t3 zo2?to{atQ}Zo3)gP2@bNLv30*+zG*9Dm41wID}ri%C>HQ0%P}Q&~457qcXA9M2ENGb!SIT!~3(S zFhhv^Jb4}??{lcDJ_k#gt03_pi@S8F4jvgO5ed^eSl4WUU9$O*A-xFp89wC{Njt-X z#a{G?p*me!=7-An?eYC-g^60=eNb_`1d;alV9ZbxUrfXvjw_Ra+q;HnRd+M>d09Ah z>G-ULu0>Jkz(nJwt#7cRua(_q zpolG!EJXVYf%euL?1G^;tkg_F!kP3LmE;~XFRK?~*rQME(04JiH#iD5HC#i@d&fb= z`6C`^tpW?ZN!YC^PHV-=q3YU9re*Vly-Rc15BtA^!k%?x?|~#7iP0vu*QSA6cbRMy75?pSr zXHs01@$#q!yhygg!@}l7a;XJG2&-eQnKF4=tw29=4q?EuBHR}@1)9~O@nWnneMX)^ zM%EVmHF-L*|0W5SeU4$hq5$0@Z^1g27IAHv#pKyqZ~SsZ41ZjW$EU-iuw+b_7}bf9 z91q0bLTb>thQ-&m{Pfvr5mxSUAuxHx@X%=qtoGy3Ur*iONZKb{RI5qHrX=C7aBy#8YaAXacd%O)~^tS@P*BtidmrEci zRs~kut5E)W7@S=DmaC{RpIqPFj1KdTqfUSV>CJu)$5OsC30j7jm8wi76$9Zz{ukKU zqC+C8>sh`nkFb1G4>Nh9PMP~;qRzf|2I2Z;!R{VUs?{`$Cciz;hO0bfH*(}BYQEB_ zq56&4pQ=dS%TN%Q01SAa*i4$a{>1pCBGDbFMAJFmxM|WJPzm=Xt@eS8=bk@UI7ge> z>RN;3y(1#v$)1NXd{x}!ZbzSJ7VJB+Y zE~fe)rKzoi6L}}agC{RK!d$9IFFOU2;KTf+M>`(h>9~Q#%C`&BYKw43#6j}Id@Xh2 z86czZp8Qzn01~gx(y?R(+Oxq1XZ<-ym%NXH6b%PbaFroFwx`IbejFos%7B#CPbMAb z<>{kR7Z9AOM3ySW(5t78;3`fencOXHJ~R6)J(l{LoUzd5oZRmqK`V&I*~fP@E}r8ZRwpkQpCh)6@1|?Acr_+xffM{iK==GrNXX9YCfGy@yj#Vz8re zKDjC}jjBp25{Z48Oi#rVyybD5vtTHQIhcQj3>9ynza9%W`~v{`_*}{Gs3{N6KimHu z03^ku_g?^@|3p1mS>Q%U`s?Ay8juk3&7>}Vq-r2n5x%KQfZ#r_Yt w?4R-f>@WZ8l*f7AzYZqk{=us8h|d4l`9GjEM;pO^AY&6U{5=2g|BLtk04h=XSpWb4 literal 0 HcmV?d00001 diff --git a/evaluate_models.sh b/evaluate_models.sh index 5f64408..a2d8952 100755 --- a/evaluate_models.sh +++ b/evaluate_models.sh @@ -13,6 +13,9 @@ python -m src.eval_main # idm python -m src.eval_main --method=idm +# GAIL +python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' + # options GAIL python -m src.eval_main --method=ogail --policy_file='checkpoints/gail-options-setobs2-Feb15_18-49-05.pt' --env='NormalizedOptionsEvalEnv' --env_kwargs='{stop_on_collision:True}' From 1624e1a349cfb50f708470c73e51bae58ca76f3c Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Thu, 17 Feb 2022 23:49:06 +0100 Subject: [PATCH 08/10] Add SHAIL --- checkpoints/sgail-options-setobs2.pt | Bin 0 -> 15019 bytes evaluate_models.sh | 7 +- src/core/gail.py | 8 +- src/core/ppo.py | 4 +- src/core/trpo.py | 8 +- src/eval_main.py | 13 ++- src/evaluation/evaluation.py | 3 +- src/safe_options/options.py | 115 ++++++++------------------- src/safe_options/policy.py | 11 ++- src/safe_options/policy_gradient.py | 4 +- src/util/wrappers.py | 4 + 11 files changed, 79 insertions(+), 98 deletions(-) create mode 100644 checkpoints/sgail-options-setobs2.pt diff --git a/checkpoints/sgail-options-setobs2.pt b/checkpoints/sgail-options-setobs2.pt new file mode 100644 index 0000000000000000000000000000000000000000..5bc3e99e565dce52f269dde8af132de183d67fc5 GIT binary patch literal 15019 zcma*O2{=~Y*FS8|JS9X_GDe|+xX)g>2XYaK>dw)FKrDY@}$x%u3FeO57jnt_W#qGbMMqBwDBWMHhDWZc&6Vyu z2QQC{>NoAUN^F$;Om~UCKPgcP?vg1{is_60c?_l4C}qzmm6(Jm)wv3N>i)|@VgvuF zRg~I)$kol~_T@@MX(U8x{>Qyngha>m)hMdJyJXBC#HwfMhZ zNIDpax5a+|_ZR~R(_{?p<+I5@gSqxI@@%cE1`g{W75uLLL622suYW03q?leY^^r}>JL;l1Th z7_KpdyvOYmB$h6P8QSw`Dm1ss~%*nmeD*tdtx zyj!mum(Opa;f{giN1uwgum3 zqtUP{W)56gvgk(SK0bBHHr}Zxm^}2e&|pm)Wf;m( zYk@6%klRoFW}k4 z=7UJrFc$+n4lqZ>pZNZoD#`61O8s8+koxcj$a5IZc8pFXwUIGotF;8)UlgF?v6d6wY0y21+~n316P-fh&$lZ1F)Idb)8s>Dey9TQcbcrbV!2 zmMa#lzf8w&AB9QtO)%iiP0&0#gZ?(k(z9nq;Lx+2%5zH50yMySq+6`6Qi0Z=M@pt_~*{Iq3+1Bo%?92zDa8h(0$Of8|RCX-vjzn@Q^T2}R z>O_|+;IFYW@JHve=Jq9+n5jgE-e(Bi?8O+qSO|&6580Cg`yey1mz+v9pwwg^6g*ae zF&>@Vus@U8;2WFZ>T^}nS*;+JlDf}dFv+3is?rd)AOma@=YakBdXoC&iz%Cn>C(%a zu=}wU+?Sdr{&hxOm|A3jEB5y=odE_UGT({AhSriOVImEFoJ)qgI=I4s6Rav<8rcX} zCOm$QEm`M}UMKE?-lBNe@zV?TnmxpbhGckHxgQLCPr+(^Nn!KC*`U8!Ae7qbPf@ey z!ulu9p!lQ&@Q@JVZ(Bi+-Yl4At4Y64yHa^vCt6<{32t)96m?_~B~+iLI@x0+JXwd{ zt#?F63Z8R!INXhJEF(5^I{^WM?P&KC2dE!nNe@2Nvb^&X#dzs;@8e#gx+7)mNvpR=Y4eK;TLOqr`&KyrE+9khQ+^R0|nhdQCx zyK&U>#|T{oZt%iVm$XZpC?Ge3QkPGGZ>dV8@9NE43MbR<;WfC6iXd#r2ORStgALYu z&Z-vn@M+8EaYHNu=-HZ~@KZ2EEGIRC)i>$;A`ZsV<^I1yR`j(qIV-WB8u^skY8-ZW-4`Hk66=6?d8a=!9nA1HWgnF9} z2y%C!9eX6j7aJ9zSVaXU>fgbDb~E^}OaWwD_rsoPuRy-87(I3OFkpQ_SxiYEENp z&e3i-dvhwj_dP&smCuM)_g@9wYjmis=OOD(R)^OHnbdfr9gI9T@@<0an=KEQfr2v! z#7B+h*QxBZg66~nG_PHaeyto08H3$0_HqS*_yCU5-b!~}c0;+}cpw@05C?gzhB`1ihzLc%?=e zw!|L6#4ZmQD3TFMWTn?hg{6_tyI;66a6LryixNItng}Y1lPOr+o5ZKnp#8NZv)6V9 zar;Q&mqH1)_4E+T?3u!650QldWzv)f8_DpF2l;)z&kV;dXKyD)z?03{IK^uMiLaev zp#__1eR49aj8x%`We(%nln1!=+IJi{pp#wAPyorKV3@s00k`RFNX8Z6jbDLK}Z|=D=g|`v#j4S8!UuTo{@(RkY+stRU9)r_CJ*0R17#}BTC2rzQ zQstx`eDWup?taN4KW_>gT=@PeaF5_EdG7dw0ZH%WPQk@mJj za4p(b)blq@yec4{DR1Z}tQ&h8%#2LLf!0qczL}@5s;g=K8&@#f7f9t6h2q|6d&tXw zG)kvk#9;9iSb0H9c>CdWN_y!oyrk_0)6PmkXh=1MYkr4<=RM@{-d^0}*a?}kt#EaY zli1ovMeJ?*h~E!Nw11S8aP*8E=sj!z->01+^2wwo6M4Cf&3L<%7x`L z(Q+OoI*xEK|2mA;XDNcITO9oSB`LOV+CVIC2o}HH^b4OX^nslPCNRc8SNy~M7Ns_8;Dy3`EW3UUhCn%7Z#alY zCaDlEmZn#844Budt0=xNPZ77H$bOVQ&6zTjO3s{OrbA})+NKq(Dclo;D`m`Ev$tob?RGTh7$&hXONnU^GW!FAxh@+ys{FGLcdx&~PEG=@M|_Em z7$*eMeab!C`-FK{Y-jf_>42`+O_mzb%hu`bWTnj>Fml;P>~VA!h$qYk!D(swCtf%^ z-Gn|J(dGviRWY4)mvF83A2uP_56;UqqLa;8&gAU_Jb2g=zKN&9<)CQb#!Q3XYnDUx z&^YF$dlU+{*fAGQhZz=~#F_hs)0^Lu*sr<;bidY+dXGqhmboK)v~esnEX{H`Mx-Y*PT+rY3S9Z<1a?h+?D;rNu!u0FSs(9^M#N8;x#v4|MW5mHJU)>BuhsN7 zZ7yt2@J8q58?fn!B*a(*uwb*vbo}0Q_*!H{>ow}=al%$C`Er&knqS@LJIat|hF+d zgjHi}V9ZZ>;lxT0PB88wG~Ubs`)Y=1ny&Euq7rzjkAZ=49kl zkc=luuWdff7{3Q5jyud>)hlCJ3oTgFrx$oEQ5l9!tYSSOO=Q`fOB0)FLGs=dxbXQ2 z&9`gf%y-nWc?n^ZwtfPg3VsI81Cv4FO%6Va@@BbL-qD?=LTXnp<{g(d1wGLt2ImMae)Zm}AET%CR=|z9$!kzdFNT z{BZhn+ma=|0J@t-fcTfDqQLX*BMcP-` z(bWF%OHl(X#~Z`;+Qa3V(4@zUA-z%w{5mh;essqL)w={2W~tu3;DpzTyl!I`*=h*5mxe%^A3(P97FHXu#58QbZ@D}Pqz9nAK>{Db!XPB1_w z1KgvEnT-Bx*4BK7PtHh&q*pK4*;iA^%XtIeaL!buzpW76v$Wtzi3DtjmBPh(KKRaI zA!i-&NKn4jhq7@b6@E{H-zCyCYr_&6aKadd?Nwuyxjy6+)#ms-AcL(aZ(;r_m)XOM z4&+|AiTydg0hK$8*qI<1{$cGZ7Sx*!as~dVXgmS`wDIf^ufy)OXpzF0{VcUag$gG) zP?hKvu3W#B|8hziLu&+Zt-um_r+lVbHVqucuEyQPop|{C9@yRN4Wm2cV7g6TzjTqJ zt<8t{&a-XYmx_(-*#w@2R!G6Yqmk6@;ZS#T`x3BGFJ*HRzGKBlKT>xx1LGN|z1wL4Tf|3+EWu`B1r5}8rnT24gx-42Y|=Sole;fr z@WSh~XC6<7EfrBqNgiU1UbDsD6=3GadW;(UfKPnzif`W_4{eExkpAmF-*NFG?{!fL z{Y*>I=4>e{g;uk%E7Mu>r2({GT^}A6_+d>#-?)$=A<$Vf3I7cF%TF=%?t4E@4&K7s z7&%Id&KS4gErk|VmsrJhF8G8EYxjxXEUm#y8wEH_ay!N)x>4-9h5X8gnM}cJ87f8~ zTi9Uu4Jy&IpljKZefa@b}*5&fz?;gO9e9<)vc z*%8{X*Ue6_%KRBVab>9gvWUV8ytuT)WZ0-~OZ%qgi#9#tpzK8e`P7Xjlls+^H+2W> z$F(3#)1|+EYq`;r2SSqm6~=!kVLIjq!0f~u2>-T+-yoVzleG^+z^h2`2Ya^rTQ2(b zjb{;uilF}T1rq?k;axyF(dXYYh_m2b(@{5A`lFCkocFBx{6 zUqRKO-*~IC3{oL;@pOSX|NHy`{+AlhcREMHpvB7(o5FF`$q)F@XBo4!X~l*a3-N(% z8Nbq8g1_Hs&fV~Fgj-*R(hQ2HVVl>|HNO!sEmj&FjNbFbJ2bIT5X$BD&0QG@nrIO6 zh`Zpn9FBjn!J@n&WMsRDH7!$SmquL?^xTpozcu-ydzq&ob+$Ts{;{I0L!sbkG6BSq z0v7*DNE2^w=LcBlvxC=HqPO58|GskzT>4bT3dav+9;dT#{=*{FX|^Kgx&HKLXa?{5 zU4wVrJd)FkRYQ%zSNOQIGwCchfc4Mb#RsRmvCeFFG`*0-9BDM`##B@+@uc$bc*?su z7!2q3XO4&W@}5CzAh%o@zS)_;_d#W7F;S2Gt=FYE%grd7H58Oz@5G2k?Z8O*3I z<~G@F!ir}GAlH-1yyB;`?Z>p4t-U4`=B{J~m*S}B)?c;&ji}9F8ow+dgUj0@4Kahq z(TS6;BgCL@83EcS<&l>$RY4)E;wxb}5 zeXm2oN`DK{wYSru=2|jM?9LAB8pw&=}M-vy&x&_gE zxauiLs%N|fe}#f22h;Y~gJERb20EM3&OP^15RM)Elev7UX4hU7v-ea_X6>Cc?^YV> zwJ(OP*A&5eeI{5e-_0!NKE;4kA*p}A$(^Zpqi@^0*?7CnT>Uj)y0D}OV?TUBs!7D$ zIz`-hG@neSW<%GiYiyePUA*MtK+omOVXo;^wrHd&?Tgw42;EfAH-IZGTUghlC*uy&fFIQd%##!o2c$7B}agP=7q<6AZcn|(tz zwgX=2=YSt=riVkvQe*2{*gr*qmPJQW$68Z3UG$CF*GP+t%MVh$**(lTqeL1{&w%6n z*N{11U1b0D63X??ql3$SVOT=}J{`G_t@GEWi%ZRD-)>vFGs2&(Hoqb$9;C>JPv)_= zPKr9UKB8sfb2i|a4Rqae!!wNw$#r9a;O>fQ_AS#2LihX7fjb%OQe+d`dtfYOONWDJ z!FYHvppceMcg51+N1TWCNH#Ea8H>~3&WU#ga)EOUKwJ|Jrf)f#9&1FxCn+E>F{F8y z%h38+5EPBAV#CL5r~Ug=MWfZ`u#tad(P4OZ{ar!fn=id3n;^;A|j zp#}|yk0jq&4{+4-b#%-0Ebp!~9OijXXUnpF@w1K1`Y`;yIqJ)YKA$j`iQUIznZYT# zF~*ZEJ-rMHZEILsK?9aaO=IbzHC#fv7M+}K2j+j>SiH_7GI%0Ohv!^p*VFraMV&6@ zS?D7wzbi*)-lRfiO*GpuZ#nHcJQl5=hC%1mMQDFEkXDfi2$SR3h0(jI<@$8G)nNo* zM&vO0b(L_j&JXrvRpNoLvAAb~1B*TSnFW^Q;la1P{M^2=PS!A!^E$8^w{n9~<@Gx> zJ9UrwM!Rbum_Mk8z2CZlKcH5J#VQk- zc7!Dr>?*=^haz^eLX~bE+KP2{X@XFZFZ(^%24`5TWrg!YIos2+tm>LNX6EW}6V-!Q zv4t%W=VZEnL{%ey$~a$j7be@f8bLLT;S{lzx(MxftMgYBDI$Vy*-#P`Q! z$XrvKoqBi;Dc~V%7-Y$QD_QXir1#T-m^3)Ao5JQ)#6e!Rh#QdG#g;_LP-^8gUZQ|e zoYYcwD|Z4U3i2^*rytjFq6-FB&m{BAWmKAYl5S4PfpZT+VX;C_-Dk;KG?|eJZhscT z`@{@#{VER@*K5J#PhX$Uy+hlLFT?Qt4v;>@7D5I#p~EF5`ZQEcsNoldsUB6N-sy_k z8Ct@F7uDg`<^AM5CWQLwNkLin0Xn?y8H;fkBQyvOfW|LbEU?RkNep{L_dnT-ErJT! zi@bMi%dt{cv(Ozrw5h`0Db=uU-+JgjJs1v3I^tPHpcl^P+4S-6u;79ruCf}y&Pl|= zysU{Z%pK_3fgLd8;d!W49Y-#_9vMCVhPxK2z?s*{;Hu*Vg|GIpim$eqs4NTdQhLIM zUCTjc>U(~elM-xuu0Z9NE;5y-eHf!~m?9N)aO{f^woiExv^ZG_9_r15rpLLUz4tn; z`5jDe)ewH!^wQx>H90tcKQN2U?UB=cl2gg5Nv-bzk^Decm1^N_o&u>UmZcq)%| z9eGT%;u%wZ;ltY{603Vk_ zG>e9eQfIRw4PpJG9jsVS0xxE5z_=ODu&rAST0acIRY&qT$#1vd3m%1vC`)ve6;Yc;c~?whL`Z!)>u(Rgk&VD+4-e9WOLEFvYyjL}DF*p2p` z@ocrLE3FJKM$=i6q7%`^eQV0)`1pP=I(w;7^X4>1)uM?|Ilz}>5?-;G;Uc`GB|+yE z)7jWo#>H6G^EzN>*;qMb>Gc!_2aVK>WE=T(g5R+B~dhO~W$SY8yp?XQUntj_*dlwU)5WwiK^S zxr}2T&%nhq?4V!fCoZsM5Y1ICN9%qHp!@a=SM5{6_4_*oss^ajW=nbc+N6VV<=`u09f1Rk@Vl9*y_r;(;Dye5MwMfWj zT7j{9%s|!E3L3*=xqw$oFnZJgs5-iyMS62|{_Ay={?&%2Qj+v=n+_(t0QNjMnVlCN z;zK8Xy>!L4_jP@US3m3>+wEJOmSwW=5-wR#Ge)KDB~Ab zKjPD_-o(&n@A-Av`|#?tQqKKIKU#BChQdXTP~YvsZ8HuK=qKAT$rJtHxRePTd@EwU z2D?o0z@^@?rPDah3Ft{FQhAX$_vBUXhu-{&trr4Zf(;pn?^bZ(AO|&+96naB+ zZ}CUG=H7_gWiR17gEpLZppz5kda=r7%Q4iZQB=3)G(X0}7@nIOLDsBV{1|wL%k<{p ztGNXnvz(24r8QA;=M-wUk))TyLLg?U9*pk-mOE)CCO3G{J2@fScI31mXGjyC?LCL3 zDejQkeuqEV)Qq~>CRnnuhApU^%L+fQWpiL1Y+Ny%Z5DiB8TPVd^W2$cKRLqIy_BJ; zL*wCvt~|^Cd6cCzD)9$8C*br+zxd-0k5Tb*JH8&8iT8BsnT1z1-_ZGqzwl-^+?rs_ z^)u-awHhYVQ{i6Tq`Cn5O)~<`I7Sl-0`Y2VDh!KCX2H)#aOyo<*pW0pd>!O0TH&}7 z*2K*QO`oHjVed%xYgag_2UWuG*ea3f*+~=|{F&dP{Djrt|IOUjoT0?$D{;rmM{G^k zYs{VTj(c=P1Q&vLvN5XG{8J?r`e?vWKieIY+M+`nLOAMAXR&6Y7b?CPOIxIVU{$g? ztzY25EQT(kx`U}~&i)VFx^pLa&52qtaO7!RvTPpFH$7VLg+tL!AGE$y$Ndf3$=kIa zXRqH42HkhMtaZvHbQty>4{FtO`Dc^alf%O~t7igCdNzm}cF96SANG6PcNL2}_mQ2R zoX_e`xxwczX83!>ID9{PJ{=v~&JE=|Fh@y-0@w6|^Gh=Y=N?O8;&^qMp}3Um4*tuA zG@oSy)mLI##9n5&(}Pl%4#2xfWsc1$70jw25`Tr>!C}{RNV1`CJ!SntoMv3j7AmWO z+YD_Sx?wE!Sej6NQvlnnWrBMmPP5-7aX9+yOU{0lKluN!g}u(4;M4mye3dCpUQh4y z7VZLp^86X>>wPE6Un5IRp$mAq8xehRSUc?_dj^KV(9`4~(v+KX@f7SfISBUl<%fR@!&?As(=-p6t{1bP=z$c|+}!hT>la}|Dle1h#XX56p1L9`%A1IT_ay00~(OPm^$@0Yg^EuY%t zHd9A$5loQ2%&s^okeXP9cBCF*z6~eo!lWkB{Wg$!o*V@k_>N9CU&kmtV+!hvlW(f~ zg5=;<>WNKZ)~Rt&c~y=Ea>wbMVI>OHjbX|X4VZ9Bk90pgf%@6gvASDZT)bb6innF4 zGY^88^wj>~`O!jXJP_-eJU6kj8`UE;XznKyGTmXxx`pcmn|CVH zy;2)kv3@G%2Tb7AhbYmoo6hL;d^1#^aAt#I#Pmyh7|_t=pg77@FnsV(GF9;=r{`tZ zWu;D?cWglRiV;k%@1%Pxl_~O5I8~JIn5=ulS0s|@fx_M9b>vypD5yh2^Gg%;K~DGbm7Ms@Qiy8Z~IS$+JafwusDsG z{>Z1gSK;Jat1r+HBnsm0n&9&n8^9nv1dTe~X=#Ka$c_Jtty(q2S28f)xKi*(F$g&b zS=xxzENb)(8o2o?y8c)XFQuI+{K`ohx2i9$J?#%$UJn&yXG_tV*hw(dd?9-5Q5J`M zSdA?$RxBpwHXFAq7A!yaKt2XhchX&uuBc+$cJ=vVFZ24o?OQl_xi5Zq%I1}q4upL7 z3H&_Ep_Jn)q*xOa==ko02DhWURi!Lzah^wJOH&0(XIAz3iqW9?yQc3OmLW{r3clvj z@Z)9!Yl$)xz;&SJVkbImauOs3^I%_h2(vsmk&YzHqz~5h5bNyT_vFSa7VWfz*K%frdqGn6@aFI_C& z%`a3282;S}e(e|oj|yDzsgVT5xn5_vA7!A{-Hbkr*v~(a{J|9mr=#U0J&NqZnx593 zXEs@3yvJ*IIu>(*6PcwjOZPnL3G9oxW1p~`f|I<`l*g!Gl>wfux7hOLO)R5VmT51} z1>HWJe6&|Ix5;)m>Tzy(B>oiJp8FlEtkjp?mB4!<0&K@o^reMV&7GsgmOdm92 zvi(oaTzd*!L|1m-PnsGU=5>O?r#v z>j=-E_k_A3(JXC>H?6q6l_`i`WAn)4n9;BmwXdub1+B~C#$R%wpsk8@bVquf=>8j4 zXttACKbb*I6Fb=A*Z?~BaU&@moX9I5K*zt9kIV32I)AB|_xt^m zziB>`JsA0wKWCyr9kpk<(OYtmlTF6=4>qH3R}*tdFogMvLulHuW$f7SQ2gV37_%(1 zSm>L{v_0Sn`L%>sYRfcesXw8&(;D7@(cesj=ApHZzWX;nhC-eTeAPw`sSMZUFI8E)^nz(t7V(6tPp`hpv&KAkSQ-K+s+8u@JCw;uw<^XHgd zb1+4{OU1_%*J8_Be^MCTi`F;OSmyZmd<&;XFU~yci>D1CeUt>gHMeFl(b4Q^oe#vi ztPrVn>wt8k70x`~_f48fl1hCY$}oE@jmu(DFI|}O)O~pF%0|)}o`z93&Y()YAM|gV z#MI?l@kXN^oSP(%*TV#~^mhuy{xo5yx@Vx?nO$_&e+0kKIGyR$sDhG|5%#nGiN)7{ zqE!4h7L+3d>F0B4%dbd!@=TYG?-Mbp8xI8HFimhQ@fAF}+{N411T*oL9Ok&Yix&hex<8D# zH$$W;Tp@XwD09h)X8zV8r2zw(tX>i|uHHb4?p*;-PG0aUJBNJd zT_ELyKiQ(7Aoi}p5O!=k33*C?QN8>e6kqj)X)71OiY+CqviAg&NSMss>+0o(UVFfD zgZJV^a$@ZpZq`(VX&b-@`fBT|QK4!h9gv5Q&F zx)V(K784CIErd&bxQ%cA4;(1j1il{!Lc*F`?B#}ry7Zp0wDq+;4xDm=+0BqVCN_S6VSSibhf5M=eBMVUdt|`vrzcn1ok~zX zj!KK(h~$RX;rFB<{)WwHD*7e`N!z4g$7fRt)QF}p72Vh>EW>Tv-Ux(MTXEcK4!XXt zVJkjb(6&-(@_(1jb8$-a&Hb?8a)}GXeA9*r6{&RNR5iE1z?p*nj3OnC<6Mu5FLCPT zlo#2z#<(_vwHs&RwmI*aP2+vmX(+<_$0KOmuG^f^)2&dD--zcG4v=QAJbyv4fcs;! zU6j}mhSKZf$WFDCoGKzoDkg^;9x#oad)7$FT{7p7Cp-Q&OqJQp2T*hK+VE%-{|I*V{|#jM&Z+?KWH zV3%#zOBPPAnPb&4=(!c$S%BUAKj9R#WKbRhWkP+IEyvTnUBqo=0= zF|IWUop%qSmSYXPfyW-kHXgwGqm%n^{cY@VZyG#UmWoMnJNY388nM?ifh5Ls*oBN4 z4i;D9sO|6AzIMEX%NqJ{umzU*a_1WiKX{6@zTXb^M=R@o>qn4o>=EWOYXqbY`OFSi z>}Q4YBf-Tx0e1F!(B4~m)HtJx^N88W-gwwD@9I=|v8_K0&l(CV28?GvWPf1tf<&;$ zyn?)jKkiKH+q+f4!QL~6*k1Pv=m8rV6k|A01Gpzqo9{n+7;Ed= zBT(I-PT`}9L3Q?YZuBo%x-Gk(W%OY~DZ=|$oqYghUNoUZ6(%c732he#WBs9KjD0m3>tu(Dn;)AB9c%^*dlp}UHOlcqf4gH4 zcWVpyoLk1U^Qwgh4_=`4$|dA245lIdkKy+k19m@ZJ{!JgiunD-eW;!yi+)ZjwCnc( zHs!FhINSI#TQ*{>(0ImWu-6KQ#t1E1x?&KvUfU|JoOPI*d|cSRyD#wf!6;DGG!af- zwhiw}4)z+sacN8owFfw+fC5P$fifNlOQFj8=5~3 zAUQun;ed@Y!r`7NkafsWSk(Ok`Y-v+^2}!8$<#RZ&Tk{z@ZB3rv*!vGsxOmQ)iJ!e zlW4Msy`LbQi+;AyC&e@&jO&8)%Mz%TxdZ6w zVTRUh^1L@j_>|iTqdZe#cl&YXqO338W0ej*dwK%#78<--VJtB}5W#@z0`-PCT=@&UXd`ENrkgXe3AZ|Ogj3RqA zgj(0%vAq(rgm*3u6&Bp|A%*XR-t)DD8*)aAr%C6C9e?@>mq!?hyB?>DU&(W1SCC01 z+k9Ea#xmGy+X04GwIN@iEdGZu=h{3Y=8Bz!#6QRXoiL{+A^0!C+<&qJCr+^c|Ie8F zcmCXH#eeR6Yxg2=sXpG|KZl3A8e{#0uJW~=M)|JwfNJMeIqle79yB}3^E|E2y9IrE?N|C~?%>rx`OuZVwA jpWQ!PPl?Hn|NKfw{6p&WaF_jua@l7iE%A^3zjXf(B+Gy^ literal 0 HcmV?d00001 diff --git a/evaluate_models.sh b/evaluate_models.sh index a2d8952..c4e5f72 100755 --- a/evaluate_models.sh +++ b/evaluate_models.sh @@ -13,6 +13,9 @@ python -m src.eval_main # idm python -m src.eval_main --method=idm +# behavior cloning +python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' + # GAIL python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' @@ -22,5 +25,5 @@ python -m src.eval_main --method=ogail --policy_file='checkpoints/gail-options-s # options GAIL-PPO python -m src.eval_main --method=ogail-ppo --policy_file='checkpoints/gail-ppo-options-setobs2-Feb15_22-05-38.pt' --env='NormalizedOptionsEvalEnv' --env_kwargs='{stop_on_collision:True}' -# behavior cloning -python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}' +# SHAIL +python -m src.eval_main --method=sgail --policy_file='checkpoints/sgail-options-setobs2.pt' --env='NormalizedSafeOptionsEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}' diff --git a/src/core/gail.py b/src/core/gail.py index d630bd4..2e34333 100644 --- a/src/core/gail.py +++ b/src/core/gail.py @@ -1,10 +1,10 @@ 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 src.core.reparam_module import ReparamPolicy +from src.core.sampling import rollout +from src.core.trpo import trpo_step +from src.core.ppo import ppo_step from tqdm import tqdm class TerminalLogger: diff --git a/src/core/ppo.py b/src/core/ppo.py index c742061..c73db2b 100644 --- a/src/core/ppo.py +++ b/src/core/ppo.py @@ -1,6 +1,6 @@ import torch -from core.sampling import rollout -from core.value_estimation import gae +from src.core.sampling import rollout +from src.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): diff --git a/src/core/trpo.py b/src/core/trpo.py index d91e363..67d0113 100644 --- a/src/core/trpo.py +++ b/src/core/trpo.py @@ -1,8 +1,8 @@ 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 +from src.core.reparam_module import ReparamPolicy +from src.core.sampling import rollout +from src.core.value_estimation import gae +from src.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): diff --git a/src/eval_main.py b/src/eval_main.py index 6b460b3..381b061 100644 --- a/src/eval_main.py +++ b/src/eval_main.py @@ -11,6 +11,8 @@ from src.evaluation.metrics import divergence, visualize_distribution from src.core.policy import SetPolicy, SetDiscretePolicy from src.core.reparam_module import ReparamPolicy from src.options import envs as options_envs2 +from src.safe_options.policy import SetMaskedDiscretePolicy +from src.safe_options import options as options_envs3 from typing import Optional, List, Dict, Tuple import torch @@ -62,8 +64,14 @@ def load_policy(method:str, policy.load_state_dict(torch.load(policy_file)) policy.eval() elif method == 'sgail': - policy = sb3.PPO.load(policy_file) - raise NotImplementedError + policy = SetMaskedDiscretePolicy(env.action_space.n) + policy( + torch.zeros(env.observation_space['observation'].shape), + torch.zeros(env.observation_space['safe_actions'].shape) + ) + policy = ReparamPolicy(policy) + policy.load_state_dict(torch.load(policy_file)) + policy.eval() else: raise NotImplementedError return policy @@ -181,6 +189,7 @@ def evaluate_policy(locations:List[Tuple[int,int]], envs_dict = dict(intersim.envs.intersimple.__dict__) envs_dict.update(dict(options_envs.__dict__)) envs_dict.update(dict(options_envs2.__dict__)) + envs_dict.update(dict(options_envs3.__dict__)) policy_metrics = [None]* len(locations) # iterate through vehicles diff --git a/src/evaluation/evaluation.py b/src/evaluation/evaluation.py index 27817ff..96b8437 100644 --- a/src/evaluation/evaluation.py +++ b/src/evaluation/evaluation.py @@ -7,6 +7,7 @@ import os import pickle from tqdm import tqdm from src.options.envs import OptionsEnv +from src.util.wrappers import OptionsTimeLimit class IntersimpleEvaluation: """ @@ -35,7 +36,7 @@ class IntersimpleEvaluation: self.env = eval_env self.n_episodes = eval_env.nv self.use_pbar = use_pbar - self.is_options_env = isinstance(self.env, OptionsEnv) + self.is_options_env = isinstance(self.env, (OptionsEnv, OptionsTimeLimit)) # metrics present on every step of every episode self.metric_keys_all = ['v_all', 'a_all', 'col_all'] diff --git a/src/safe_options/options.py b/src/safe_options/options.py index a8d0e50..eeb5bae 100644 --- a/src/safe_options/options.py +++ b/src/safe_options/options.py @@ -3,14 +3,18 @@ import numpy as np import torch from stable_baselines3.common.vec_env import DummyVecEnv as VecEnv -from core.reparam_module import ReparamPolicy +from src.core.reparam_module import ReparamPolicy from tqdm import tqdm -from core.gail import train_discriminator, roll_buffer, TerminalLogger +from src.core.gail import train_discriminator, roll_buffer, TerminalLogger from dataclasses import dataclass -from safe_options.policy_gradient import trpo_step, ppo_step +from src.safe_options.policy_gradient import trpo_step, ppo_step import torch.nn.functional as F -from safe_options.collisions import feasible +from src.options.envs import OptionsEnv +from src.safe_options.collisions import feasible + +from intersim.envs import IntersimpleLidarFlatIncrementingAgent +from src.util.wrappers import OptionsTimeLimit, Setobs, TransformObservation @dataclass class Buffer: @@ -156,81 +160,6 @@ def rollout(env_fn, policy, n_episodes, max_steps_per_episode): return (states, safe_actions, actions, rewards, dones), (ll_states, ll_actions, ll_rewards, ll_dones) -class OptionsEnv(gym.Wrapper): - - def __init__(self, env, options): - super().__init__(env) - self.ll_action_space = env.action_space - self.options = options - self.action_space = gym.spaces.Discrete(len(options)) - self.max_plan_length = max(t for _, t in options) - - def plan(self, option): - target_v, t = option - current_v = self.env._env.state[self.env._agent, 1].item() - dt = self.env._env._dt - a = (target_v - current_v) / (t * dt) - a = self.env._normalize(a) - a = a * np.ones((t,)) - a += 0.01 * np.random.randn(*a.shape) - a = np.clip(a, self.ll_action_space.low, self.ll_action_space.high) - return a - - def execute_plan(self, obs, option, render_mode=None): - observations = np.zeros((self.max_plan_length + 1, *self.env.observation_space.shape)) - actions = np.zeros((self.max_plan_length + 1, *self.ll_action_space.shape)) - rewards = np.zeros((self.max_plan_length + 1,)) - env_done = np.ones((self.max_plan_length + 1,), dtype=bool) - plan_done = np.ones((self.max_plan_length + 1,), dtype=bool) - infos = [] - - plan = self.plan(option) - observations[0] = obs - env_done[0] = False - for k, u in enumerate(plan): - plan_done[k] = False - o, r, d, i = self.env.step(u) - actions[k] = u - rewards[k] = r - env_done[k+1] = d - infos.append(i) - observations[k+1] = o - - if render_mode is not None: - self.env.render(render_mode) - - if d: - break - - n_steps = k + 1 - return observations, actions, rewards, env_done, plan_done, infos, n_steps - - def step(self, action, render_mode=None): - a = int(action) - assert a == action - ll_obs, ll_actions, ll_rewards, ll_env_done, ll_plan_done, ll_infos, ll_steps = self.execute_plan(self.last_obs, self.options[a], render_mode) - hl_obs = ll_obs[ll_steps] - hl_reward = (ll_rewards * ~ll_plan_done).sum().item() - hl_done = ll_env_done[ll_steps].item() - hl_infos = { - 'll': { - 'observations': ll_obs, - 'actions': ll_actions, - 'rewards': ll_rewards, - 'env_done': ll_env_done, - 'plan_done': ll_plan_done, - 'infos': ll_infos, - 'steps': ll_steps, - } - } - self.last_obs = hl_obs - return hl_obs, hl_reward, hl_done, hl_infos - - def reset(self, *args, **kwargs): - self.last_obs = super().reset(*args, **kwargs) - return self.last_obs - - class SafeOptionsEnv(OptionsEnv): def __init__(self, env, options, safe_actions_collision_method=None, abort_unsafe_collision_method=None): @@ -287,7 +216,7 @@ class SafeOptionsEnv(OptionsEnv): o, r, d, i = self.env.step(u) actions[k] = u rewards[k] = r - env_done[k+1] = d + env_done[k] = d infos.append(i) observations[k+1] = o @@ -303,3 +232,29 @@ class SafeOptionsEnv(OptionsEnv): n_steps = k + 1 return observations, actions, rewards, env_done, plan_done, infos, n_steps + +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 NormalizedSafeOptionsEvalEnv(max_episode_steps=float('inf'), safe_actions_collision_method=None, abort_unsafe_collision_method=None, **kwargs): + return OptionsTimeLimit(SafeOptionsEnv(Setobs( + TransformObservation(IntersimpleLidarFlatIncrementingAgent( + n_rays=5, + **kwargs, + ), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10)) + ), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5)], safe_actions_collision_method=safe_actions_collision_method, abort_unsafe_collision_method=abort_unsafe_collision_method), max_episode_steps=max_episode_steps) diff --git a/src/safe_options/policy.py b/src/safe_options/policy.py index 4e0e9d0..0325f59 100644 --- a/src/safe_options/policy.py +++ b/src/safe_options/policy.py @@ -2,7 +2,7 @@ import torch import torch.nn as nn from torch.distributions import Categorical from torch.distributions.kl import kl_divergence -from core.policy import SetDiscretePolicy +from src.core.policy import SetDiscretePolicy class SetMaskedDiscretePolicy(SetDiscretePolicy): @@ -21,6 +21,15 @@ class SetMaskedDiscretePolicy(SetDiscretePolicy): a = super().torch_dist(logits).probs return (a * (1 - z)).sum(-1) + def predict(self, observations, state=None, episode_start=None, deterministic=True): + observation = torch.tensor(observations['observation']) + safe_actions = torch.tensor(observations['safe_actions']) + if deterministic: + _, actions = self.forward(observation, safe_actions).max(-1) + else: + actions = self.sample(self.forward(observation, safe_actions)) + return actions, None + # def torch_dist_nomask(self, dist): # print('no mask logprob') # logits = dist[..., :self.action_dim] diff --git a/src/safe_options/policy_gradient.py b/src/safe_options/policy_gradient.py index 6e968e4..803af15 100644 --- a/src/safe_options/policy_gradient.py +++ b/src/safe_options/policy_gradient.py @@ -1,6 +1,6 @@ import torch -from core.value_estimation import gae -from core.optimization import conjugate_gradient, line_search +from src.core.value_estimation import gae +from src.core.optimization import conjugate_gradient, line_search def trpo_step(value, policy, states, safe_actions, actions, rewards, dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1): diff --git a/src/util/wrappers.py b/src/util/wrappers.py index 3916e75..3a088f9 100644 --- a/src/util/wrappers.py +++ b/src/util/wrappers.py @@ -9,6 +9,10 @@ class TransformObservation(gym.wrappers.TransformObservation): def __getattr__(self, name): return getattr(self.env, name) +class OptionsTimeLimit(gym.wrappers.TimeLimit): + def __getattr__(self, name): + return getattr(self.env, name) + class CollisionPenaltyWrapper(Wrapper): def __init__(self, env, collision_distance, collision_penalty, *args, **kwargs): From 84351e77f28a84a9683597c246996e9208932bfd Mon Sep 17 00:00:00 2001 From: ebuehrle <43623224+ebuehrle@users.noreply.github.com> Date: Fri, 18 Feb 2022 06:54:52 +0100 Subject: [PATCH 09/10] Add SHAIL-PPO --- .../sgail-ppo-options-setobs2-17-02-2022.pt | Bin 0 -> 16279 bytes evaluate_models.sh | 3 +++ src/eval_main.py | 4 ++++ src/safe_options/options.py | 2 +- 4 files changed, 8 insertions(+), 1 deletion(-) create mode 100644 checkpoints/sgail-ppo-options-setobs2-17-02-2022.pt diff --git a/checkpoints/sgail-ppo-options-setobs2-17-02-2022.pt b/checkpoints/sgail-ppo-options-setobs2-17-02-2022.pt new file mode 100644 index 0000000000000000000000000000000000000000..33df960884d6f3b421d3499c80a2c22cf667288d GIT binary patch literal 16279 zcmbum2{@MB*EejQ$B+z(sHDuKxc1tHQqmwzN`sW56j!98NoJxllqu094HBi|+G{tU zkDb44ocr3pwb!}NW$(4u{w*(0NhvWgIXSWa zO;DGtzK_Sb+6eccT5U@O8 z!N{Qb;R-^r!QMh~Qy;}1(aV-uSXrzL2wD^vE|i!b9?~P%V(yCYpx`itxyu9mR|Eww zm>V9jEG%UC+z1;5p`@v)k8BSttf%F?MFB#ozaa|l!6Ea)C;wfkAe8nI?`dNtl<|@5 zxz7(>A(S2OBM~l?o8&3xY2&FBFO>IG@)VC3DkKH|)z_BW`R9##{~^Lgs2wOf zzD7P=*ynFWY~qDFJ^J)0qWiZTy>Owvx3KR&ir85ErHH{FQf!6&{#L~BzZ5b0tw{gB z1sI15O}vEz{!zs4A4LZKA;L~Ls7H}G;^9Km-!gh$D_%JGZ$-@hmNO(=IMiEc{*NMd ze=B0~hZK9EWsf4V;X*@P<@%rD| zzlxaVANFr)AMRoeG{B43N*R&o^7zk zb1U{sm&0vY)i9uP5X)Aw=V;rG#!9}uvd;;9_(Zl0w zJxKYA)2J-mf-ZwqFh|M{f)yz)6Ao1pz8sn8Mp<(AOrDzg1lf>$;O+I8nw+AMHVDwFV{^Z(`+!3OGx411|if zMrSQeXBQ@~fOQuIyiSl$++t*ilVr8%{rGKoHm(VFjcXQ#rsUzlrg%ZdmHnVS|1rCH z^cKdqB=HjmOW?zRQf4+~3Z`C}jR`f3><2M8n`j7Kxzos4_L+3f>Le!ZyRqOcrSbKE zgDo1twJ{M4Z|oy!d-PH1n=Onv#mK8{Rg!spHmu~uEK3@@5kJSzM0dSef?rd&;Fm&CRmF-oc-wI- zIJTb$n|oqJ-cp?&-TH`Cr}f3=2W1!>#YAXI0ozwJ5A>Kghy-!?@(~O7i?r8oZye zki0cjCa-TDB@K%c@pNBtT)Q>IYMl&83!()(Uu=XgCe`@h zQX#rcTm&=3N=ZiZe$nNIGPIErlIc3xqV8U9_|6~!UbK}|kz;_amllHLw%5=nDiuC= zRkAt5Qi%S79kBGBKX(Y&2D#sstsgNskkV}{t|vnC%r;NA8BzH@{Qliec@HgfV1tmXm- z)6?LF!8zQ1DwbtC9VF#R*`l(0rO;XY0<{bD>5JrBP+HQ5WXp=djDReGV@4_F>db|c zL*`-Dxlw3+#|Z|W-iprtL$RVhg^zpIA)4Nh0Ut^;M75>u)g{GSVL6|Kt|xS1=!Yp_ zx}gN4w6yV1R1l6jI)(1`Xu`viLNF40MrIjiLP*_3JiR)NZCsWPn~iP=4s^z%M{f~s z82bfYjvB|!Kedx9tyL)5M+Ws3`oXLXQnWp^hMD$8>@zJ6_de|>2-u(kbBD#V&l6&} z@yu+P78VHN$=ArSl}ZrxDTBGpdmwP>a|z{Ab1><{QV1=sB6nA(q4ZTB>hwJYk16H~ z$iYEi78egnjT-pEHJNotY$1_rJehJ-2z#Tu80SsD0=r|!^3vxO?A%~8?3l0(TNbZ@ zE(bp_{BRgnj5tQt#}}de(j0c)WFRQa-it?95_EQ!bf*VHpY?(PDtF@Zh(~L6F_NJrIw%yW6G9nJ% zzqyQd&r)%GzjN$Yh$bu>6$_`o*x=>Yli94)VmRLW4fEUko&50M1CNTYVbX#ZWLeKz zq2G%M=KdT&HIh6a?NSNe9b`}3o}X}{ZWF$w>ESAbQQGjs^S>&Xpm|gO0SomBeJh4A8^75jN9(O z>!Aag@zglx{%9}pu8zWvy?zj~^)#@^^Ki~coKBhc1atfO@asw5Y_fS4Y@hvzg!Zf( zN50L+kG+RO<6&LqeZ`HGbef?b4~A7k?QqG)H8f#M8vL9dgC|Q$aQ67s?3Uh1mLHS^ zKGB+@wPUxDvKj7Z88sEor6qvHbT4r9vxF~Y_sQEuW?qK-rJ2TLqb5K#U5(Xt6BMiR4`qBiadD0P{QOIDyX_bqW%HW zX%LQ9%Gp%LG>1HkaE9mK3sB?pdhBv|f^R+-iu&ZnquJ}Da3ABLlUad#q%nr>h(Wcv za`d9B8=4lC!7R^_ERm@~zauxGpGGz_S$T!^k^RZ$&gzA)Oe=7j)Fyr}J_Qb{SP_*O zy#;j_jbNk80vM^W0$q)#;sQYt+DU(}_INQKzK?!@HEz26NXSQ29JLabDNSUkGJh<`KS$G$#JBry4M1jF-{c-I^)UeXi}#_x8(juDK^s2%}7a~nks zHtMufu!THW+=Z`(XyHKr7)U+3nZK?1j5faI4k7+aw0ao zzsj_8J`0>wj*#E}+R^Jlu6{jE`|t8U`D-SB`fJwiu7>}MzxK!GrRQPxsRLwpgg1Ov zKO$JW@26<<#xHQmax_TROW@+o1cy9a1Ycc7vRYAO%YYc@cC z(*#`BtV-X@=#o2Udidp9a#L9TSQ{)GK-rS7oEaIio89;c+QpXOZX@t}W~|7n~Ae~y!lyX)}(&vE+q zxc&3&KSLV&E{cYoWfr6?V>~{nk*2p}&l1Ogov0+cj5POghXXf~Sm$nYJUh^UhWRfR zJdryov5_=?i7Wo5+3f!KAtD2FB%AVbo_G z@V%qUrJNtZ-tr9AaCJVoTvz~4=7?eQqgl}Q?jZEJy@=nK(ZKpgKg3qWd*uDDFv!d8 z5Vfr@RaPh1w4jFy~$38E25_>d>>c%bvKJ6$?SxErXj)}yMjb^tF8$tW= z3Vd{~3a4rhAmhjN#kpT~`N$9Q45rIbl*|?#xVHt;)?0{r?eXCYLL#xzOoQe0k7e#V z7UJ&6>v-{=0Uk1543;yWqf4YdeW^TzR_qYq7>B1U{mvX6Q7qeuMG55n z_jUMvusGFRRm8>{kB2+Qngs6;-Xi@U9K-a?A@p{u6kHxSm`>|Iixw7!Ld#1PXzY$< zh1EataaT6FDsF(CVP<@<>1!Akl}@zZ#}lO!ta|hN-h#`&j)P-oA$&K|=C?!qN#QYF z4B9Y>G|YNUPa|~-&TR>Ib#KPDq+bA&^N~OlE z1IH&%1g&><_zd|}P&XX}hZR)uL%lgooUKbkChODD!3Mm3gCR6)ZU>V$Cj<#<=@6`W zo+a8nM)hiQS}L|1lx*~9jJ-9TakG$I@xbLb*wVsA7)5U#ka|1SBfq-P>}_-@{@S>d?eBSyaPt;TLU*U%sAY$XGU?O z>BP;QxN4aKR4NRjWP1+i^}9-Vmp(1+c3@o##=^ufQzz{{AMp09pCDD!1fx5ZvCv+M zXODUVkM28ylI=^ZS=~fDjO_(8zKkNf^_rY6TODFvQQlBU%qB|d zRQt&Uh>A-LcqlDn%85ZRZfdOHW1tf6NF0Jw&PmWSC(GEGOm)~{_yty%WVRQ6e;E$wMxE>b+@dlUhsPR|Uc3%Q6XqQw+k6DP?<9~<}D)NZmKo>r)`xP!7 zJ(atikcKIDB&okv5qK>K;t4JlXwV(VwdV@?tGO&<-?5|`pH{co5=(hpQoGBUe zD~J}<)-t~*C&~OzsnBgTkOd7cVr308x$%*Q`L$`C8*kV3~P{n!d;nBO$FGihL1c~8E zNoU^HHJ{&n94;7iQW}Rlx08Z`qp&}v4GzxD=$ShN)LUUVf0kE>`guq3QoO4`DP|Dc zrW%gRo0aj1?@O@w_Js6#SqXJbW_YN56kMOIjNK`7=sktquLqi2O~=WbN{OEE6^wi}0#uu02=7KPNv#4ya}$iOkR_#+saUx}f(;z6gD2HQ zsBo_ZR!ooR={g+W$7Dd3v^6~tiU4}wvW-Ow-l z8$ONPirrl)#Axbu931P7Z);m09 zBPyma6&228;-)Ay7^j?zdtdJc$Bx73KTeB38)gVL4%$#M@I6bZe9c^Y$FP8*ed*Gx zLj^ayBC+$l4lJ8<8HUL9$dKSTHu)wqm~Tbs<{0cZz!xX07Zcx$Lzv;UUf?@PjSrD} zLaOg^xTq-0hucZ%wzj$U4l>{StcHCx0^mH-nO&W-l*5d@V9blE! zhr5ZiVZdw-wtalC!SS|W|CJGN?p!an#KnRtO`agy`EdvOj!J;FtHWTnLD2p;Rd7wD~AL+)HMpf4}l6HR_rbfO@qx++?NXKMF{+4C;KuCCq~_@fj`i(U)D zhaLfsZPQ@x-3U^DN0I00DAI*jI$%M26MOUIFuA6l1mD6dL3~Ow=6q6xGv*gr{SG}G zuSRI#MKQSQ6N3`lvPkpyOem_?rsEgx2b#2nStr}jcFUfv=j7{qGu?`|XDqUkz5>O(8BTeA!lPbt8?CtFc& zK?t)-BQ!W=FuPh9K|)+@>7M9DxU=m(E0%DB*n?*1dm{%T&3X~zAM$+qc}Xgf(T~ct zH8S?$DfZI;4AvXwL7`=To_49nj@R5}=Z<|8jjR3zXQC|OB}>39K6!%pmjfaHls^5w zMTHKVzn!hvWlzVu8gQxBiF~7K5&1Yj8D@7TGF_8*#Qou6+?{An@5TEgF5HCEL#M)B zJ7SMl`XLIi;FW!-W-W+>-MwE4vCTFK&gG zMlYC4yfy4C+ee1yWWd?lv*_JV8teNVh6Ntc?14IG(Ob_mx6f{{OsEdp?@HOg^OVYr z-VXM0y{L+#GqEE_a9dY8hUvP|mi{L}#(4`>Jh_y<_SA>jYu^bJM`WT`NEwd3XULT= zDbSppNIpDZGL84zMZQ^OgZG+ESl*^hst>#;hfCJLzPq((rxVB;)61FbI#;k!odIVL zeI$WRGx^{jqhQL_`{b!*8jkQ^Np`kfU^`3Z()X*&p?z2y?RfDDlKZYAhpgVR5eIw5 zR7;lEo!pBty^n%Mmn-}#`++}4?;wqfn!)?f1T4aC8&P*Ds;L}h9TAsY+C;=^!jMZkLmP- zhvxz2TD)QwO1t6uI~6$5J4WE%5|69wr0DA2;h-|70f*{D39Q=&vZuMrASsfPGZuzu zFxmqStB>Mm?`rYrhhE&YmmH}ls&t%XCu|6q!{7Sl;-oh*yuWrD&z_o&OQQ79cH?@I zG)adRr3|I^RoeWQS|UC4Ae^uFZ)G>fN#f%ZVI+%&lXHg+*oq}4$cI%zy!truvojsm zXt?3G;tLq3+agH*zL+M<%3)Sw5D6^U3;98#z^kZ*ojBpfJqC;MoufAL<~k`FQ+yH{ z&y2y(MRxp<*eftl>xnNaj>n{Fcfn`QHJEN}{{=@w>c4ErB|8&FrtD zlRKoj`<;C-``J}U-u?krHPphl7-{r1452r7>heSDqIt!18G1_H0FE3fM7y+jIJQ-t z?<<=}x4d3Pqhb{J-g#Z%I$|?Z*6;C4&&~o*y{SC3SQf-j#=_bU`LL_48apwJ&9cy^ z0Z~SLR(oIIPZ!Y91NOtM;6gI8P>h;49|8J&DVeY&nq_KT!P(s};90d4U7od=%Ks>H{uJ84U-IFmTn^qYpfHLc+dD zEW&UB?eKBs@0^l(`T0gzYcUzdUI{@AYa!#jTk)%s1sz>7m0p!om-|Jp7KWR{D6{ZUDV^`Y`U;cb16TThIf6=`c8@o_)J} z3(W4A&}RkFa6sJ!%fbs;jPzEF@I210WgdZ;#kV0|rx+(LTMosNH&}r6WE>Gi;H^w} zmExijqH-w_D_%~+M{}+-k2VYDbNr29#VJ`dnwuq>d&5BVCg3O;6}OAk)I7t3RemJG z%U6)xPz7rPN0SoO<1j$X79MO@cWN}Pf?LgoOjY6|+pnffG{%aNkc9)_{mw^N6zGOe zz2*4Tz+$}kWjApR(dN1mJHR+nnvQrZ58ur@Sh~0=XtqUQ;GU^aoU6_2@12MH^H0Os ziN{b+!c_Ee(H8h>Bn~$lO=)-BTHMpH7eQ=0M6J0FBOOOGS+{I#8L10hIbvw^VjenM zcc9+X*X)hSOe`v5s3CH~xEN1Nc|42?Y)%j*4_SUfObxHzw&990ov5p+&KH&6gK6Rp zSSwe7>}e#kYn2x~q+Z1Cd_I+{8^-4a>GAc~J-FkevAlcUS&`wGTbS-S1Nz=@p^*i_ z+(M}qD?iDzJpn^`uW%!(m+ZkW1n2Q>ZA$cq;{mjd-p7I+%As!SWxI5nKj8Nt1#D=I4n5p40$P9O zFx6e3QNunDHIuaXzMLzh%X%F;`X`f`4covt-cHmaZ_E`cYhh#e3v|pI3U7WbAxCQ? zSn8JuSXnIvQWi4!e8zGjcW4`V-WnpfnV^7%R^3?JyaS(F93rs|ThTpI3j*wR2$p`I zC+OJP#9qW7!aBWPqTy9itinSJopR;z>$R7nyX&0kf;;8db>SrHv`%4;kJ2F4SCP+E z?!vRFbzo&-b4sCOzi4UL6Vc?rZ*21Yr)=B1!7!pG3+;!yDJ`1}xq-)#Uz)k{c^m$BS2nsm_ZrTDauHSyiQN%Ved z78q3N(#HqhS1(b_N8tfviJG0@Jk%>1p7g8hLmGGd)Kd+OzYS%pnx(Bj?@rfNs5Vscqo?_74leFx@SUhj#SO7rUV zqoL48gEkb^!Mee!ob|noB_Az$+=;=wbk01C5YHydg506;d2f32SU!1o{Rg%MKZUa* z7aqCY2p?rBp_1z`g!5;~v6q}p5N^U@HKyRxb%1;rI2B b%jH3h+~3e{?7tLK6ay z;!zb{j2+p%5491EdyY4q(V zEoveaL*s>K;ZXNh_Hm9py6m?C#oRvJ=V=ly+Y$$cISTYt@&()+HUKwB4xqDs++wz- zMUdqiO=K*dW42=|_WSjLc;qKmvs4>Yec*yA%`5O!??a-6^Tc?Vwj_Ogv={vt`Uz(! zj77h3m4x`p!ktrvpj}|a)z9^X^uynXV*D+pWoOJ5JzN1#y)~JSu@PPBV}wf#53z@> zn;_Hg1yOIA2eShjuyM|5oO9j~q%>^U*SVJvyp7 z2~E#li4*H(`99Ync<^*ON&I|_DD6*R$&HoZ_l*$;M|XIHmYC<2Nmfn#if(DHw87RJ zT21EA@h+2k=4xG0md`cj)^`vDOBLea(;c|4u}xU{piZ4TgV0F{qS&U2~)LD z0jbbE(9gnxJk`IAewSmI*MyCDb8mlU6M95s?w(G}UT@&L!ewy71Q$5+{62O~2tv)? zPphgWdei=|T_C9JAoZV_&f*;8@Z8omGD>Y36TK1R!bEeF49$X1*CTNHoCH)SDxmUB z2d0cL=Az?nf*eb8{5igZO%8m;mQ|`y71s&;W8wtV$REp-3^(v$w|bG5q1OcgV=IVM z-DPO@IRQgY?!iM+lj-`^Cj8nkeWt3Pic2(Y*r$dBa%ARm{5WIG*0~KSsnl<+){H28>V`3QmtK$cTgPe4R7HuoC1Adks+A;sB97 z>C7spZNa0js^Md6Z~Aa+U$m_3LiaC8@N%sxoI8McD4@wH_-r@Jb6X+kf1{3cKi)%L z&M3r?fiVzNjj%JxgbuATq|tGCeK|CDD3&d5| zpFxV*Ubz3`EY`gd!o^kjM7NJDHs#8p>0D>-c1;;}gpGu#dLy>%Mh$*?luEXY>qF0s z*b9BHItuLXZ-v>Vl~8!xi+$Hz3+cgw@c85qyrdZb_hkma9oUb7YaX+`h2zNO!#7B` z+ZIvwnIUkZF%6nd{=hZHKKN;h43F*~O~hvSfPr!?YUFL#sa@GIu5V* z>{-TCu3!NwU0^eKB3+bo2;1a*X~U5v+)867`?XX9pFR+g@g3h8S(}LMVG}U9x)lzZ zmkX{c*+8G^H`o`}2m2&nCO;lbqI=$}pyifvSfc zw%ZAM)lw$<;skeF86KVLhI*{_$PSZoP2ZY8+4!I!w{d7`&l z9vicL6fD==L~q_{M>{7gzIsL$e^!vmYhTY59JQ&(rWZS)Mc13oDP`=vwmPp<3I*Ax z`uxT2&G7QjR~GQ`60UE00501s;LWULTv2)*9QGuNs_6^po*s(6vHDzZPY3H>szz^j z%i*QU&1hTx0>z0c9Gq|-uQodHtjW%B+G`@tl$%HIrf30lk7N%&#bQ9;A=nUMMv|8( z!?zD1wC3R%5Gs9z<>g1Q@3Uk${P7IDOMSuOeN*9r!UuNu0ms?(k3lZO`7P1{AJ-W3J8P z1;a-+!r~EzsI@v0_xYy~+FpdCer|(;7iLuLvpZfjI0&~sorf8xs-Qum82esJfa+Js z20RajC+52$x51S>UgE%Gdg9Yo21Dt+*CIM5BbK_o)I^QFr4Z&&hRN<@`SxOY$f_uW z=^ZU>={yhmNX-ZpPgaq<^20E2&p{Lq%*E#;9*Oq+D!{SBU!u5;0pC@(mGsgKftM$o z@Y;P_{#_s)OD1@9#00P(uLi@<4-gf!c<^;?2GlLv znEJXNr;bJGbhv~biyQkKRMW+IvbF`!9&?8EIhY5_f80W;FS^vBehbP?ih`Ws!N4OV zc=_gQ?B}T{c=EuM@0>al*Y}*))PItI;n&{4@Wu#y7;VZY2iU@e(c6Fxu7TpD$!NDI z9Q}+< z<9NP*^i)`XkwV;Yb(pQ6C}^zNBfx7*x#NVfv^}mj9hTDx!_)*+e78OvUO16Dteb=U zhCauaYp9QI3;3)laL|DP6WXp4syZ`)Z3$wQBkd9$XVN|hR1(Jp7} zNjYh$zXFDX4nyet6ih!JEi#@UPRG{{g_4*`c(pqjj1L&Gd1mo=As`4QmWXp1xWe|m zX&@n?`Q+-7TDVgZ0NVp=(O|DJx~(5bi_^y7jgnmKdfYBr^XwE0O|c+$O`2#e(1K&C zW102*7^stNp7_V{CN-DY-2J)8s$lCzstQqrHQT8u<;sed|xyhyjZem$fRft^cV!M8n(fv>l;x_rvM|1 zzLA>hp1AB*f7GorhAW}>S;~VTa`B|6&etaR?ni{k`RRQUX|7899{vCyu3{}F-(8Eh^kpI6 z{~a?aR-{sQOu6Gccah8B5}5Zu6?M&x@QL1RC~R@U9aS;3M9U6WUam#0{(dl3*$cXR z&Ttat6?v(1Ec(B&5G^U#gFCOv^RGj~fM$j>wMEz9rn4a(^*vtH-+Cn9*?A84Bu9$I z`B~Gkqmr>&&5dbvR)h77Ls*hkN7VY9CY=}6+2fQaaBj@b>T|<71O~}_p}u09pmgn2 zQGTByRL&j%@omkp=)xQ{zo$ey?UZWS$o|Ko+%P*5&|eF43Ikw#VkEoo7=ybF_v6ZPc~~cPQ}jl! ziacCWM-0EtfmKn<*$tZ|IPsw+bNZA|RFdzaMxH5VHUz<)!Zkhdy_+~`*Ky`wxE_wq zlIIx?(Rg&72KH7^!<_k%=#s5RH@_3`!&k54x3C=kb@XmN&Ts}>^LZoc?h=PrGoM4C z!ASUcKNdGWx1w|VyKuL|n~1d2Ewr9+mgR+hB^lK%W@9tlp9Rj;&g z$HYDOalbUYH;$#3P4#f!sVcOxi6Gzfnn+}80c)ur$R}MJhGubFxUJhbl-}MSU2RRs zw}HywH?|h%gi7M*10rbO^@AnPX+=-_a2EO{2A0PT<7p#$diIX^nEB# z+jRVRvTQLLeOtu~e+c=IuF4jK2zRQL^?Dva2n70}4Z(#T=QUk8$yN1dmYm zhmYpDkl}QU-OV1#hfHr^D_$h=&YD&hr2YdWAKD0XZ-2sFT`F|$=mE50%Ud+FN@mx- z%Fy{vjgS*CgvW>8gu)Y!*zA-D?4c{{y1fySG+g-c79aXn-GG+MRZ|kUiFSWECOFy? zNA60xfvv*+d}{G6w%4ryoIk7b9|y!~VMzhz-_C%2k4i|tF`d|E8YilBOa_(h`uyF! z=j^J>XIxYJjx2Y|0fi?UG0iWC7t9<2+8X6hxZfQXhe=?e-#hZkYaC5b9>ms_mczHv z^GWNTP4HE$3Jk)3jh|#(C)b z#uEiKbI?~!oo9VqM=!t6Co;Qy==1ha8viXBCOAsMqQl{W4C!X<^|(Lx(S8c$`zCT< zzYSDs$XxjNEE2MI%|(Zp`J&a%x{%_n184fD;ev(T?EE6Y5Br0`EM=15_2cs_ck3`- zJjNPkt(Xmmnp07@G6<7g&$6kS(e(A|cW^8(m!7eB0bdg?gURJ2qV68ROIfB!zdi{{ zdw(TshL^#u6H#j$}(hW1wW!2ylzL#iF(aV8Om-yq~+0=IADa!NOtm zwVXX2qK5|^?|4stAO?M6p>-JFRc;S@~w_-&>_76TC`ViHLb1S zC~1JEzeWqTcI-p_N@aX+pN7Zg8!p9AkT z{^tHu)baG<1wUFvP1a}O=7t7X7B!B)a7so0q*yMm9?zGRyW{w!BSbC1{v=bl6XI?X zy65EwjQtrxuUy)|pE%euhkHA*b&ns>_^b|R-)+R+A*)ewpfT)zU!f*BdyYW)(K52OfE%3&rJyZ4hV#ThMswBHEAZh;ikzi zd?l5GANrVb_x|%yvLn1IqQ#znS<(hyw{4;B$DiTekQAKp&>EzU^~DRU3J%qK^4JPV zzS>ril9YQWS*wC$&dE_KrHz^~X4E3q3#DS0;*(JsxN~z8yt<_i!$nHmIZ}^Ktm_c1 z3o+&sT{>Csicg|}ZGHK?@Ld8SACJQJIQDDKYj`X64J!wG^4s5b!i3^dQgtpB*SyO{ z`9@XR5)q3xhUr4zbA@<2aR;Y|6NEv^y9!=RQW?8=!E(MoR%>UAt0+7fDE z#*&F5$+ENPax(>j*)00$+bqnjPQ{_md*c_y8L(#S6ufw4BERdjo#dE*#-{75$s)gl z_(SEJK>o~OaPKnV57rs;jHV$#yN1!`U=il+al^5TO=9r2;t2RhX?05N}czAYy^6+fkU9J9q^Y9KuXTo;ZTv5%G!5~}xguETR zjODD~iLXs(qDgp~)1b0^vTeQ}IyR>=uQ88Uq=GFR9G^x~qXNL^x*|K|pMule(n$i# z5}4gc#LUn*w)3lrz&d>!({^hW=`7uiuTbWl#FfBzKDS=B=WuUjZI#^^5B7?dC`9Nc~ zQ92pF#%QtoALjqX^At}Qr15{iU$pzb;V+h~RN_~MXQEB--rVw$8XfyBf{10Gg^rGH z*!n5O$!mBZdkr5oG$KvNhI&L z9hLT8gx9IhL|UB^)x%C*Av)(9NeW8?{cF$Ryi`vvN|_y>*5F8_#Rt%mE?sV;d;xY% zpH2k%QW$y30QBWY^Dp;)kdsFuXymXX(3$#}ybP3Nxrw^e`Qtsf*qx7_OHBE*#ti;4 zGKlV%x1lHAOd`weZFsGwierDTXqsNLp1xk^!hQ80p+HrP=NBBn#%(h(!)YcZn;UBrT!K;ynK(fsb1iJ^*GpnZ2ooiwt>6kq)^j6`g zYY66kjK!hX;vjo<7E7{N!$vI(#ckHnJhh9WYKa9mNXvn&MKNGq6bn&R2cUMxBn&Tf zqSD%n(5#1-HO^5$1rzmXeEd69-ZP9h&ufHXZ89R)SUKu`KApsi_s8T~#PG5_OuduL zq9Sg=D9u6qwQ&OWZMQpxRHnN839tjg8&YT?Il2XxzKi8EDyfs#cgh^gmM=XV=WT`?O@ zziP$mX}Ub7J%QTYD1;d|#=yY11kkr02-V@c@p0O8+_=aIuc)PAgRUy=7M{oT+wQ_T zO=k?U^}ud?icDe*iM$krj<=Vf_R@vStvZXrNHS&mxV_@5 zlVMst?wgZ>YTjb>cyb-d8d)XE-*uSH&k_oT3G{j4`=^k(-woBQUV~@#N#beJv)|9$ zjFk`lF;Ds#E0?&9*vwJ)mKfK!{|L1|eZg7nlE8Gr8XWRD6szy+!FR20keRoc#Uy*- z{hf}uu-8P`^V|nC7bwwzX9h!3XCPEAHw3vep|JPa4)_FZP-piY>Fnwrj(j%OM|}pv zycE{4vj*d*)#KaXLnu6xjlLBjXe4(S*3XX>kRH3bx~T&7Ixgd`*RJq+wiU)-NnxY6 z-4ncbVQ^vQj-EJNPwdt%o=pJ_ zT$Q9h5y9|2MtFC?DRyCF9#IZIj6J$wcJ5tN3X+G) zxJFc$93*P@-vIJc>(S+zHMJ~Az{@{V1Tm}n(AJurXfeDC9|r34#&{WCb#Ols)83A6 zUWkFYy<4@+t+jNM;zsH>Y&f4-I|$9DU1gi^^~E_;Dxl#Rp}UJz;0F#z`$v8FcBA`{ zQ#+ckiaN;@<+7Pe#~11w6H7HJ^KqU;B6_cU4r>qOk?yH$q3o<1p7h$o3RE8Q(qehR zt;7#BA|(Pph|ZyD4@dc?Wj1tO1OL@@3|8rH=kJaC(p|^n z*?vI-HPJ1C#BL|i;k~!9)sbVz_&BEA-W&2)>f^yPMX+<+c+|Dir{f+gP-%(xNE6QU z=f|{2=JeUXa{x|kNeV9M7ca{HzwfQ&pKiBa;v3~ad!rJ~D`=2)P zCpLT7zp!@y#{TEG{u7%h_!rjx-`M~Bf~tRv^UMFh{+DfddCJHP{jWv_lIH*2{GTlP rf9~zS^V54SZ}RuO2lqRxUCf^T{rvl9FHdQi-%oNqf4|%R-S+ Date: Fri, 18 Feb 2022 10:18:10 +0100 Subject: [PATCH 10/10] Fix imports --- .../trpo/experiments/sgail-options-setobs2.py | 19 +++++++++++-------- .../experiments/sgail-ppo-options-setobs2.py | 15 +++++++++------ 2 files changed, 20 insertions(+), 14 deletions(-) diff --git a/scratch/etienne/trpo/experiments/sgail-options-setobs2.py b/scratch/etienne/trpo/experiments/sgail-options-setobs2.py index 02c8bff..3c43ed5 100644 --- a/scratch/etienne/trpo/experiments/sgail-options-setobs2.py +++ b/scratch/etienne/trpo/experiments/sgail-options-setobs2.py @@ -1,19 +1,22 @@ # %% +import sys +sys.path.append('../../../../') + 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 +from src.safe_options.options import gail +from src.core.gail import Buffer +from src.core.value import SetValue +from src.safe_options.policy import SetMaskedDiscretePolicy +from src.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 +from src.util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np -from safe_options.options import SafeOptionsEnv +from src.safe_options.options import SafeOptionsEnv from torch.utils.tensorboard import SummaryWriter -from core.reparam_module import ReparamPolicy +from src.core.reparam_module import ReparamPolicy obs_min = np.array([ [-1000, -1000, 0, -np.pi, -1e-1, 0.], diff --git a/scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py b/scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py index 5365a94..a7b01e7 100644 --- a/scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py +++ b/scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py @@ -1,16 +1,19 @@ # %% +import sys +sys.path.append('../../../../') + 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 +from src.safe_options.options import gail_ppo, Buffer +from src.core.value import SetValue +from src.safe_options.policy import SetMaskedDiscretePolicy +from src.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 +from src.util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs import numpy as np -from safe_options.options import SafeOptionsEnv +from src.safe_options.options import SafeOptionsEnv from torch.utils.tensorboard import SummaryWriter obs_min = np.array([