Port TRPO, PPO, GAIL
This commit is contained in:
74
scratch/etienne/trpo/core/discriminator.py
Normal file
74
scratch/etienne/trpo/core/discriminator.py
Normal file
@@ -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)
|
||||
125
scratch/etienne/trpo/core/gail.py
Normal file
125
scratch/etienne/trpo/core/gail.py
Normal file
@@ -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
|
||||
39
scratch/etienne/trpo/core/optimization.py
Normal file
39
scratch/etienne/trpo/core/optimization.py
Normal file
@@ -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
|
||||
96
scratch/etienne/trpo/core/policy.py
Normal file
96
scratch/etienne/trpo/core/policy.py
Normal file
@@ -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))
|
||||
72
scratch/etienne/trpo/core/ppo.py
Normal file
72
scratch/etienne/trpo/core/ppo.py
Normal file
@@ -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
|
||||
162
scratch/etienne/trpo/core/reparam_module.py
Normal file
162
scratch/etienne/trpo/core/reparam_module.py
Normal file
@@ -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)
|
||||
73
scratch/etienne/trpo/core/sampling.py
Normal file
73
scratch/etienne/trpo/core/sampling.py
Normal file
@@ -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
|
||||
23
scratch/etienne/trpo/core/test_optimization.py
Normal file
23
scratch/etienne/trpo/core/test_optimization.py
Normal file
@@ -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)
|
||||
79
scratch/etienne/trpo/core/trpo.py
Normal file
79
scratch/etienne/trpo/core/trpo.py
Normal file
@@ -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
|
||||
48
scratch/etienne/trpo/core/value.py
Normal file
48
scratch/etienne/trpo/core/value.py
Normal file
@@ -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)
|
||||
40
scratch/etienne/trpo/core/value_estimation.py
Normal file
40
scratch/etienne/trpo/core/value_estimation.py
Normal file
@@ -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
|
||||
Reference in New Issue
Block a user