Paper code basis
This commit is contained in:
@@ -0,0 +1,270 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class DistributionalQNetwork(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
n_obs: int,
|
||||
n_act: int,
|
||||
num_atoms: int,
|
||||
v_min: float,
|
||||
v_max: float,
|
||||
hidden_dim: int,
|
||||
device: torch.device = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.net = nn.Sequential(
|
||||
nn.Linear(n_obs + n_act, hidden_dim, device=device),
|
||||
nn.ReLU(),
|
||||
nn.Linear(hidden_dim, hidden_dim // 2, device=device),
|
||||
nn.ReLU(),
|
||||
nn.Linear(hidden_dim // 2, hidden_dim // 4, device=device),
|
||||
nn.ReLU(),
|
||||
nn.Linear(hidden_dim // 4, num_atoms, device=device),
|
||||
)
|
||||
self.v_min = v_min
|
||||
self.v_max = v_max
|
||||
self.num_atoms = num_atoms
|
||||
|
||||
def forward(self, obs: torch.Tensor, actions: torch.Tensor) -> torch.Tensor:
|
||||
x = torch.cat([obs, actions], 1)
|
||||
x = self.net(x)
|
||||
return x
|
||||
|
||||
def projection(
|
||||
self,
|
||||
obs: torch.Tensor,
|
||||
actions: torch.Tensor,
|
||||
rewards: torch.Tensor,
|
||||
bootstrap: torch.Tensor,
|
||||
discount: float,
|
||||
q_support: torch.Tensor,
|
||||
device: torch.device,
|
||||
) -> torch.Tensor:
|
||||
delta_z = (self.v_max - self.v_min) / (self.num_atoms - 1)
|
||||
batch_size = rewards.shape[0]
|
||||
|
||||
target_z = (
|
||||
rewards.unsqueeze(1)
|
||||
+ bootstrap.unsqueeze(1) * discount.unsqueeze(1) * q_support
|
||||
)
|
||||
target_z = target_z.clamp(self.v_min, self.v_max)
|
||||
b = (target_z - self.v_min) / delta_z
|
||||
l = torch.floor(b).long()
|
||||
u = torch.ceil(b).long()
|
||||
|
||||
l_mask = torch.logical_and((u > 0), (l == u))
|
||||
u_mask = torch.logical_and((l < (self.num_atoms - 1)), (l == u))
|
||||
|
||||
l = torch.where(l_mask, l - 1, l)
|
||||
u = torch.where(u_mask, u + 1, u)
|
||||
|
||||
next_dist = F.softmax(self.forward(obs, actions), dim=1)
|
||||
proj_dist = torch.zeros_like(next_dist)
|
||||
offset = (
|
||||
torch.linspace(
|
||||
0, (batch_size - 1) * self.num_atoms, batch_size, device=device
|
||||
)
|
||||
.unsqueeze(1)
|
||||
.expand(batch_size, self.num_atoms)
|
||||
.long()
|
||||
)
|
||||
proj_dist.view(-1).index_add_(
|
||||
0, (l + offset).view(-1), (next_dist * (u.float() - b)).view(-1)
|
||||
)
|
||||
proj_dist.view(-1).index_add_(
|
||||
0, (u + offset).view(-1), (next_dist * (b - l.float())).view(-1)
|
||||
)
|
||||
return proj_dist
|
||||
|
||||
|
||||
class Critic(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
n_obs: int,
|
||||
n_act: int,
|
||||
num_atoms: int,
|
||||
v_min: float,
|
||||
v_max: float,
|
||||
hidden_dim: int,
|
||||
device: torch.device = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.qnet1 = DistributionalQNetwork(
|
||||
n_obs=n_obs,
|
||||
n_act=n_act,
|
||||
num_atoms=num_atoms,
|
||||
v_min=v_min,
|
||||
v_max=v_max,
|
||||
hidden_dim=hidden_dim,
|
||||
device=device,
|
||||
)
|
||||
self.qnet2 = DistributionalQNetwork(
|
||||
n_obs=n_obs,
|
||||
n_act=n_act,
|
||||
num_atoms=num_atoms,
|
||||
v_min=v_min,
|
||||
v_max=v_max,
|
||||
hidden_dim=hidden_dim,
|
||||
device=device,
|
||||
)
|
||||
|
||||
self.register_buffer(
|
||||
"q_support", torch.linspace(v_min, v_max, num_atoms, device=device)
|
||||
)
|
||||
self.device = device
|
||||
|
||||
def forward(self, obs: torch.Tensor, actions: torch.Tensor) -> torch.Tensor:
|
||||
return self.qnet1(obs, actions), self.qnet2(obs, actions)
|
||||
|
||||
def projection(
|
||||
self,
|
||||
obs: torch.Tensor,
|
||||
actions: torch.Tensor,
|
||||
rewards: torch.Tensor,
|
||||
bootstrap: torch.Tensor,
|
||||
discount: float,
|
||||
) -> torch.Tensor:
|
||||
"""Projection operation that includes q_support directly"""
|
||||
q1_proj = self.qnet1.projection(
|
||||
obs,
|
||||
actions,
|
||||
rewards,
|
||||
bootstrap,
|
||||
discount,
|
||||
self.q_support,
|
||||
self.q_support.device,
|
||||
)
|
||||
q2_proj = self.qnet2.projection(
|
||||
obs,
|
||||
actions,
|
||||
rewards,
|
||||
bootstrap,
|
||||
discount,
|
||||
self.q_support,
|
||||
self.q_support.device,
|
||||
)
|
||||
return q1_proj, q2_proj
|
||||
|
||||
def get_value(self, probs: torch.Tensor) -> torch.Tensor:
|
||||
"""Calculate value from logits using support"""
|
||||
return torch.sum(probs * self.q_support, dim=1)
|
||||
|
||||
|
||||
class Actor(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
n_obs: int,
|
||||
n_act: int,
|
||||
num_envs: int,
|
||||
init_scale: float,
|
||||
hidden_dim: int,
|
||||
std_min: float = 0.05,
|
||||
std_max: float = 0.8,
|
||||
device: torch.device = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.n_act = n_act
|
||||
self.net = nn.Sequential(
|
||||
nn.Linear(n_obs, hidden_dim, device=device),
|
||||
nn.ReLU(),
|
||||
nn.Linear(hidden_dim, hidden_dim // 2, device=device),
|
||||
nn.ReLU(),
|
||||
nn.Linear(hidden_dim // 2, hidden_dim // 4, device=device),
|
||||
nn.ReLU(),
|
||||
)
|
||||
self.fc_mu = nn.Sequential(
|
||||
nn.Linear(hidden_dim // 4, n_act, device=device),
|
||||
nn.Tanh(),
|
||||
)
|
||||
nn.init.normal_(self.fc_mu[0].weight, 0.0, init_scale)
|
||||
nn.init.constant_(self.fc_mu[0].bias, 0.0)
|
||||
|
||||
noise_scales = (
|
||||
torch.rand(num_envs, 1, device=device) * (std_max - std_min) + std_min
|
||||
)
|
||||
self.register_buffer("noise_scales", noise_scales)
|
||||
|
||||
self.register_buffer("std_min", torch.as_tensor(std_min, device=device))
|
||||
self.register_buffer("std_max", torch.as_tensor(std_max, device=device))
|
||||
self.n_envs = num_envs
|
||||
self.device = device
|
||||
|
||||
def forward(self, obs: torch.Tensor) -> torch.Tensor:
|
||||
x = obs
|
||||
x = self.net(x)
|
||||
action = self.fc_mu(x)
|
||||
return action
|
||||
|
||||
def explore(
|
||||
self, obs: torch.Tensor, dones: torch.Tensor = None, deterministic: bool = False
|
||||
) -> torch.Tensor:
|
||||
# If dones is provided, resample noise for environments that are done
|
||||
if dones is not None and dones.sum() > 0:
|
||||
# Generate new noise scales for done environments (one per environment)
|
||||
new_scales = (
|
||||
torch.rand(self.n_envs, 1, device=obs.device)
|
||||
* (self.std_max - self.std_min)
|
||||
+ self.std_min
|
||||
)
|
||||
|
||||
# Update only the noise scales for environments that are done
|
||||
dones_view = dones.view(-1, 1) > 0
|
||||
self.noise_scales = torch.where(dones_view, new_scales, self.noise_scales)
|
||||
|
||||
act = self(obs)
|
||||
if deterministic:
|
||||
return act
|
||||
|
||||
noise = torch.randn_like(act) * self.noise_scales
|
||||
return act + noise
|
||||
|
||||
|
||||
class MultiTaskActor(Actor):
|
||||
def __init__(self, num_tasks: int, task_embedding_dim: int, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.num_tasks = num_tasks
|
||||
self.task_embedding_dim = task_embedding_dim
|
||||
self.task_embedding = nn.Embedding(
|
||||
num_tasks, task_embedding_dim, max_norm=1.0, device=self.device
|
||||
)
|
||||
|
||||
def forward(self, obs: torch.Tensor) -> torch.Tensor:
|
||||
task_ids_one_hot = obs[..., -self.num_tasks :]
|
||||
task_indices = torch.argmax(task_ids_one_hot, dim=1)
|
||||
task_embeddings = self.task_embedding(task_indices)
|
||||
obs = torch.cat([obs[..., : -self.num_tasks], task_embeddings], dim=-1)
|
||||
return super().forward(obs)
|
||||
|
||||
|
||||
class MultiTaskCritic(Critic):
|
||||
def __init__(self, num_tasks: int, task_embedding_dim: int, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.num_tasks = num_tasks
|
||||
self.task_embedding_dim = task_embedding_dim
|
||||
self.task_embedding = nn.Embedding(
|
||||
num_tasks, task_embedding_dim, max_norm=1.0, device=self.device
|
||||
)
|
||||
|
||||
def forward(self, obs: torch.Tensor, actions: torch.Tensor) -> torch.Tensor:
|
||||
task_ids_one_hot = obs[..., -self.num_tasks :]
|
||||
task_indices = torch.argmax(task_ids_one_hot, dim=1)
|
||||
task_embeddings = self.task_embedding(task_indices)
|
||||
obs = torch.cat([obs[..., : -self.num_tasks], task_embeddings], dim=-1)
|
||||
return super().forward(obs, actions)
|
||||
|
||||
def projection(
|
||||
self,
|
||||
obs: torch.Tensor,
|
||||
actions: torch.Tensor,
|
||||
rewards: torch.Tensor,
|
||||
bootstrap: torch.Tensor,
|
||||
discount: float,
|
||||
) -> torch.Tensor:
|
||||
task_ids_one_hot = obs[..., -self.num_tasks :]
|
||||
task_indices = torch.argmax(task_ids_one_hot, dim=1)
|
||||
task_embeddings = self.task_embedding(task_indices)
|
||||
obs = torch.cat([obs[..., : -self.num_tasks], task_embeddings], dim=-1)
|
||||
return super().projection(obs, actions, rewards, bootstrap, discount)
|
||||
@@ -0,0 +1,424 @@
|
||||
import math
|
||||
from typing import Sequence, Union
|
||||
|
||||
import distrax
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
from flax import nnx
|
||||
|
||||
from reppo.jaxrl import utils
|
||||
|
||||
|
||||
def torch_he_uniform(
|
||||
in_axis: Union[int, Sequence[int]] = -2,
|
||||
out_axis: Union[int, Sequence[int]] = -1,
|
||||
batch_axis: Sequence[int] = (),
|
||||
dtype=jnp.float_,
|
||||
):
|
||||
"TODO: push to jax"
|
||||
return nnx.initializers.variance_scaling(
|
||||
0.3333,
|
||||
"fan_in",
|
||||
"uniform",
|
||||
in_axis=in_axis,
|
||||
out_axis=out_axis,
|
||||
batch_axis=batch_axis,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
|
||||
class UnitBallNorm(nnx.Module):
|
||||
def __call__(self, x: jax.Array) -> jax.Array:
|
||||
return x / (jnp.linalg.norm(x, axis=-1, keepdims=True) + 1e-8)
|
||||
|
||||
|
||||
def normed_activation_layer(
|
||||
rngs, in_features, out_features, use_norm=True, activation=nnx.swish
|
||||
):
|
||||
layers = [
|
||||
nnx.Linear(
|
||||
in_features=in_features,
|
||||
out_features=out_features,
|
||||
kernel_init=torch_he_uniform(),
|
||||
rngs=rngs,
|
||||
)
|
||||
]
|
||||
if use_norm:
|
||||
layers.append(nnx.RMSNorm(out_features, rngs=rngs))
|
||||
if activation is not None:
|
||||
layers.append(activation)
|
||||
return nnx.Sequential(*layers)
|
||||
|
||||
|
||||
class Identity(nnx.Module):
|
||||
def __call__(self, x: jax.Array) -> jax.Array:
|
||||
return x
|
||||
|
||||
|
||||
class FCNN(nnx.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
hidden_dim: int = 512,
|
||||
hidden_activation=nnx.swish,
|
||||
output_activation=None,
|
||||
use_norm: bool = True,
|
||||
use_output_norm: bool = False,
|
||||
layers: int = 2,
|
||||
input_activation: bool = False,
|
||||
*,
|
||||
rngs: nnx.Rngs,
|
||||
):
|
||||
if layers == 1:
|
||||
self.module = normed_activation_layer(
|
||||
rngs,
|
||||
in_features,
|
||||
out_features,
|
||||
use_norm=use_output_norm,
|
||||
activation=output_activation,
|
||||
)
|
||||
else:
|
||||
if input_activation:
|
||||
input_layer = nnx.Sequential(
|
||||
# nnx.LayerNorm(in_features, rngs=rngs) if use_norm else Identity(),
|
||||
hidden_activation,
|
||||
normed_activation_layer(
|
||||
rngs,
|
||||
in_features,
|
||||
hidden_dim,
|
||||
use_norm=use_norm,
|
||||
activation=hidden_activation,
|
||||
),
|
||||
)
|
||||
else:
|
||||
input_layer = nnx.Sequential(
|
||||
normed_activation_layer(
|
||||
rngs,
|
||||
in_features,
|
||||
hidden_dim,
|
||||
use_norm=use_norm,
|
||||
activation=hidden_activation,
|
||||
)
|
||||
)
|
||||
hidden_layers = [
|
||||
normed_activation_layer(
|
||||
rngs,
|
||||
hidden_dim,
|
||||
hidden_dim,
|
||||
use_norm=use_norm,
|
||||
activation=hidden_activation,
|
||||
)
|
||||
for _ in range(layers - 2)
|
||||
]
|
||||
output_layer = normed_activation_layer(
|
||||
rngs,
|
||||
hidden_dim,
|
||||
out_features,
|
||||
use_norm=use_output_norm,
|
||||
activation=output_activation,
|
||||
)
|
||||
self.module = nnx.Sequential(
|
||||
input_layer,
|
||||
*hidden_layers,
|
||||
output_layer,
|
||||
)
|
||||
|
||||
def __call__(self, x: jax.Array) -> jax.Array:
|
||||
return self.module(x)
|
||||
|
||||
|
||||
class CriticNetwork(nnx.Module):
|
||||
def __init__(
|
||||
self,
|
||||
obs_dim: int,
|
||||
action_dim: int,
|
||||
hidden_dim: int = 512,
|
||||
use_norm: bool = True,
|
||||
use_encoder_norm: bool = False,
|
||||
use_simplical_embedding: bool = False,
|
||||
encoder_layers: int = 1,
|
||||
head_layers: int = 1,
|
||||
pred_layers: int = 1,
|
||||
*,
|
||||
rngs: nnx.Rngs,
|
||||
):
|
||||
self.feature_module = FCNN(
|
||||
in_features=obs_dim + action_dim,
|
||||
out_features=hidden_dim,
|
||||
hidden_dim=hidden_dim,
|
||||
hidden_activation=nnx.swish,
|
||||
output_activation=utils.multi_softmax if use_simplical_embedding else None,
|
||||
use_norm=use_norm,
|
||||
use_output_norm=use_encoder_norm,
|
||||
layers=encoder_layers,
|
||||
rngs=rngs,
|
||||
)
|
||||
self.critic_module = FCNN(
|
||||
in_features=hidden_dim,
|
||||
out_features=1,
|
||||
hidden_dim=hidden_dim,
|
||||
hidden_activation=nnx.swish,
|
||||
output_activation=None,
|
||||
use_norm=use_norm,
|
||||
use_output_norm=False,
|
||||
layers=head_layers,
|
||||
rngs=rngs,
|
||||
)
|
||||
self.pred_module = FCNN(
|
||||
in_features=hidden_dim,
|
||||
out_features=hidden_dim,
|
||||
hidden_dim=hidden_dim,
|
||||
hidden_activation=nnx.swish,
|
||||
output_activation=utils.multi_softmax if use_simplical_embedding else None,
|
||||
use_norm=use_norm,
|
||||
use_output_norm=False,
|
||||
layers=pred_layers,
|
||||
rngs=rngs,
|
||||
)
|
||||
|
||||
def features(self, obs: jax.Array, action: jax.Array):
|
||||
state = jnp.concatenate([obs, action], axis=-1)
|
||||
return self.feature_module(state)
|
||||
|
||||
def critic_head(self, features: jax.Array) -> jax.Array:
|
||||
return self.critic_module(features)
|
||||
|
||||
def critic(self, obs: jax.Array, action: jax.Array) -> jax.Array:
|
||||
features = self.features(obs, action)
|
||||
return self.critic_head(features)
|
||||
|
||||
def critic_cat(self, obs: jax.Array, action: jax.Array) -> jax.Array:
|
||||
features = self.features(obs, action)
|
||||
return self.critic_head(features)
|
||||
|
||||
def forward(self, obs, action):
|
||||
features = self.features(obs, action)
|
||||
value = self.critic_head(features)
|
||||
return self.pred_module(features), value
|
||||
|
||||
|
||||
class CategoricalCriticNetwork(nnx.Module):
|
||||
def __init__(
|
||||
self,
|
||||
obs_dim: int,
|
||||
action_dim: int,
|
||||
hidden_dim: int = 512,
|
||||
use_norm: bool = True,
|
||||
use_encoder_norm: bool = False,
|
||||
use_simplical_embedding: bool = False,
|
||||
encoder_layers: int = 1,
|
||||
head_layers: int = 1,
|
||||
pred_layers: int = 1,
|
||||
num_bins: int = 51,
|
||||
vmin: float = -10.0,
|
||||
vmax: float = 10.0,
|
||||
*,
|
||||
rngs: nnx.Rngs,
|
||||
):
|
||||
self.num_bins = num_bins
|
||||
self.vmin = vmin
|
||||
self.vmax = vmax
|
||||
|
||||
self.feature_module = FCNN(
|
||||
in_features=obs_dim + action_dim,
|
||||
out_features=hidden_dim,
|
||||
hidden_dim=hidden_dim,
|
||||
hidden_activation=nnx.swish,
|
||||
output_activation=utils.multi_softmax if use_simplical_embedding else None,
|
||||
use_norm=use_norm,
|
||||
use_output_norm=use_encoder_norm,
|
||||
layers=encoder_layers,
|
||||
rngs=rngs,
|
||||
)
|
||||
self.critic_module = FCNN(
|
||||
in_features=hidden_dim,
|
||||
out_features=self.num_bins,
|
||||
hidden_dim=hidden_dim,
|
||||
hidden_activation=nnx.swish,
|
||||
output_activation=None,
|
||||
use_norm=use_norm,
|
||||
use_output_norm=False,
|
||||
layers=head_layers,
|
||||
input_activation=not use_simplical_embedding,
|
||||
rngs=rngs,
|
||||
)
|
||||
self.pred_module = FCNN(
|
||||
in_features=hidden_dim,
|
||||
out_features=hidden_dim,
|
||||
hidden_dim=hidden_dim,
|
||||
hidden_activation=nnx.swish,
|
||||
output_activation=None,
|
||||
use_norm=use_norm,
|
||||
use_output_norm=False,
|
||||
layers=pred_layers,
|
||||
input_activation=not use_simplical_embedding,
|
||||
rngs=rngs,
|
||||
)
|
||||
self.zero_dist = nnx.Param(
|
||||
utils.hl_gauss(jnp.zeros((1,)), num_bins, vmin, vmax)
|
||||
)
|
||||
|
||||
def features(self, obs: jax.Array, action: jax.Array):
|
||||
state = jnp.concatenate([obs, action], axis=-1)
|
||||
return self.feature_module(state)
|
||||
|
||||
def critic_head(self, features: jax.Array) -> jax.Array:
|
||||
cat = self.critic_module(features) # + self.zero_dist.value * 40.0
|
||||
return cat
|
||||
|
||||
def critic_cat(self, obs: jax.Array, action: jax.Array) -> jax.Array:
|
||||
features = self.features(obs, action)
|
||||
return self.critic_head(features)
|
||||
|
||||
def critic(self, obs: jax.Array, action: jax.Array) -> jax.Array:
|
||||
value_cat = jax.nn.softmax(self.critic_cat(obs, action), axis=-1)
|
||||
value = value_cat.dot(
|
||||
jnp.linspace(self.vmin, self.vmax, self.num_bins, endpoint=True)
|
||||
)
|
||||
return value
|
||||
|
||||
def forward(self, obs, action):
|
||||
features = self.features(obs, action)
|
||||
value_cat = jax.nn.softmax(self.critic_head(features), axis=-1)
|
||||
value = value_cat.dot(
|
||||
jnp.linspace(self.vmin, self.vmax, self.num_bins, endpoint=True)
|
||||
)
|
||||
return self.pred_module(features), value
|
||||
|
||||
def __call__(self, obs: jax.Array, action: jax.Array) -> jax.Array:
|
||||
features = self.features(obs, action)
|
||||
value_cat = jax.nn.softmax(self.critic_head(features), axis=-1)
|
||||
value = value_cat.dot(
|
||||
jnp.linspace(self.vmin, self.vmax, self.num_bins, endpoint=True)
|
||||
)
|
||||
pred = self.pred_module(features)
|
||||
return value, value_cat, pred
|
||||
|
||||
|
||||
class SACActorNetworks(nnx.Module):
|
||||
def __init__(
|
||||
self,
|
||||
obs_dim: int,
|
||||
action_dim: int,
|
||||
hidden_dim: int = 512,
|
||||
ent_start: float = 0.1,
|
||||
kl_start: float = 0.1,
|
||||
use_norm: bool = True,
|
||||
layers: int = 2,
|
||||
min_std: float = 0.1,
|
||||
*,
|
||||
rngs: nnx.Rngs,
|
||||
):
|
||||
self.actor_module = FCNN(
|
||||
in_features=obs_dim,
|
||||
out_features=action_dim * 2,
|
||||
hidden_dim=hidden_dim,
|
||||
hidden_activation=nnx.swish,
|
||||
output_activation=None,
|
||||
use_norm=use_norm,
|
||||
use_output_norm=False,
|
||||
layers=layers,
|
||||
input_activation=False,
|
||||
rngs=rngs,
|
||||
)
|
||||
start_value = math.log(ent_start)
|
||||
kl_start_value = math.log(kl_start)
|
||||
self.temperature_log_param = nnx.Param(jnp.ones(1) * start_value)
|
||||
self.lagrangian_log_param = nnx.Param(jnp.ones(1) * kl_start_value)
|
||||
self.min_std = min_std
|
||||
|
||||
def actor(
|
||||
self, obs: jax.Array, scale: float | jax.Array = 1.0
|
||||
) -> distrax.Distribution:
|
||||
loc = self.actor_module(obs)
|
||||
loc, log_std = jnp.split(loc, 2, axis=-1)
|
||||
std = (jnp.exp(log_std) + self.min_std) * scale
|
||||
pi = distrax.Transformed(distrax.Normal(loc=loc, scale=std), distrax.Tanh())
|
||||
return pi
|
||||
|
||||
def det_action(self, obs: jax.Array) -> jax.Array:
|
||||
loc = self.actor_module(obs)
|
||||
loc, _ = jnp.split(loc, 2, axis=-1)
|
||||
return jnp.tanh(loc)
|
||||
|
||||
def temperature(self) -> jax.Array:
|
||||
return jnp.exp(self.temperature_log_param.value)
|
||||
|
||||
def lagrangian(self) -> jax.Array:
|
||||
return jnp.exp(self.lagrangian_log_param.value)
|
||||
|
||||
def __call__(self, obs: jax.Array) -> jax.Array:
|
||||
loc = self.actor_module(obs)
|
||||
loc, std = jnp.split(loc, 2, axis=-1)
|
||||
return jnp.tanh(loc), std, self.temperature(), self.lagrangian()
|
||||
|
||||
|
||||
class TD3ActorNetworks(nnx.Module):
|
||||
def __init__(
|
||||
self,
|
||||
obs_dim: int,
|
||||
action_dim: int,
|
||||
hidden_dim: int = 512,
|
||||
ent_start: float = 0.1,
|
||||
kl_start: float = 0.1,
|
||||
use_norm: bool = True,
|
||||
layers: int = 2,
|
||||
min_std: float = 0.1,
|
||||
*,
|
||||
rngs: nnx.Rngs,
|
||||
):
|
||||
self.actor_module = FCNN(
|
||||
in_features=obs_dim,
|
||||
out_features=action_dim * 2,
|
||||
hidden_dim=hidden_dim,
|
||||
hidden_activation=nnx.swish,
|
||||
output_activation=None,
|
||||
use_norm=use_norm,
|
||||
use_output_norm=False,
|
||||
layers=layers,
|
||||
input_activation=False,
|
||||
rngs=rngs,
|
||||
)
|
||||
start_value = math.log(ent_start)
|
||||
kl_start_value = math.log(kl_start)
|
||||
self.temperature_log_param = nnx.Param(jnp.ones(1) * start_value)
|
||||
self.lagrangian_log_param = nnx.Param(jnp.ones(1) * kl_start_value)
|
||||
self.min_std = min_std
|
||||
|
||||
def actor(
|
||||
self, obs: jax.Array, scale: float | jax.Array = 1.0
|
||||
) -> distrax.Distribution:
|
||||
loc = self.actor_module(obs)
|
||||
loc, log_std = jnp.split(loc, 2, axis=-1)
|
||||
std = (jnp.exp(log_std) + self.min_std) * scale
|
||||
pi = distrax.Transformed(distrax.Normal(loc=loc, scale=std), distrax.Tanh())
|
||||
return pi
|
||||
|
||||
def det_action(self, obs: jax.Array) -> jax.Array:
|
||||
loc = self.actor_module(obs)
|
||||
loc, _ = jnp.split(loc, 2, axis=-1)
|
||||
return jnp.tanh(loc)
|
||||
|
||||
def temperature(self) -> jax.Array:
|
||||
return jnp.exp(self.temperature_log_param.value)
|
||||
|
||||
def lagrangian(self) -> jax.Array:
|
||||
return jnp.exp(self.lagrangian_log_param.value)
|
||||
|
||||
|
||||
class TD3DeterministicDist(distrax.Distribution):
|
||||
def __init__(self, loc: jax.Array, scale: float | jax.Array):
|
||||
self.loc = loc
|
||||
self.scale = scale
|
||||
|
||||
def sample(self, seed=None):
|
||||
return self.loc + self.scale * jax.random.normal(seed, self.loc.shape)
|
||||
|
||||
def log_prob(self, value: jax.Array) -> jax.Array:
|
||||
return jnp.zeros_like(value)
|
||||
|
||||
def sample_and_log_prob(self, *, seed, sample_shape=...):
|
||||
sample = self.sample(seed=seed)
|
||||
log_prob = self.log_prob(sample)
|
||||
return sample, log_prob
|
||||
@@ -0,0 +1,374 @@
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.distributions import constraints
|
||||
from torch.distributions.transforms import Transform
|
||||
from torch.distributions.normal import Normal
|
||||
|
||||
from reppo.torchrl.reppo import hl_gauss
|
||||
|
||||
|
||||
class TanhTransform(Transform):
|
||||
r"""
|
||||
Transform via the mapping :math:`y = \tanh(x)`.
|
||||
|
||||
It is equivalent to
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
ComposeTransform(
|
||||
[
|
||||
AffineTransform(0.0, 2.0),
|
||||
SigmoidTransform(),
|
||||
AffineTransform(-1.0, 2.0),
|
||||
]
|
||||
)
|
||||
|
||||
However this might not be numerically stable, thus it is recommended to use `TanhTransform`
|
||||
instead.
|
||||
|
||||
Note that one should use `cache_size=1` when it comes to `NaN/Inf` values.
|
||||
|
||||
"""
|
||||
|
||||
domain = constraints.real
|
||||
codomain = constraints.interval(-1.0, 1.0)
|
||||
bijective = True
|
||||
sign = +1
|
||||
log2 = torch.log(torch.tensor(2.0)).to("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
def __eq__(self, other):
|
||||
return isinstance(other, TanhTransform)
|
||||
|
||||
def _call(self, x):
|
||||
return x.tanh()
|
||||
|
||||
def _inverse(self, y):
|
||||
# We do not clamp to the boundary here as it may degrade the performance of certain algorithms.
|
||||
# one should use `cache_size=1` instead
|
||||
return torch.atanh(y)
|
||||
|
||||
def log_abs_det_jacobian(self, x, y):
|
||||
# We use a formula that is more numerically stable, see details in the following link
|
||||
# https://github.com/tensorflow/probability/blob/master/tensorflow_probability/python/bijectors/tanh.py#L69-L80
|
||||
return 2.0 * (self.log2 - x - torch.nn.functional.softplus(-2.0 * x))
|
||||
|
||||
|
||||
def get_activation(name):
|
||||
if name == "gelu":
|
||||
return nn.GELU()
|
||||
elif name == "relu":
|
||||
return nn.ReLU()
|
||||
elif name == "swish":
|
||||
return nn.SiLU()
|
||||
elif name is None:
|
||||
return nn.Identity()
|
||||
else:
|
||||
raise ValueError(f"Unknown activation: {name}")
|
||||
|
||||
|
||||
def normed_activation_layer(
|
||||
in_features, out_features, use_norm=True, activation="swish", device=None
|
||||
):
|
||||
layers = [nn.Linear(in_features, out_features, device=device)]
|
||||
if use_norm:
|
||||
layers.append(nn.RMSNorm([out_features], device=device))
|
||||
if activation is not None:
|
||||
layers.append(get_activation(activation))
|
||||
return nn.Sequential(*layers)
|
||||
|
||||
|
||||
class FCNN(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_features,
|
||||
out_features,
|
||||
hidden_dim=256,
|
||||
hidden_activation="swish",
|
||||
output_activation=None,
|
||||
use_norm=True,
|
||||
use_output_norm=False,
|
||||
layers=2,
|
||||
input_activation=False,
|
||||
device=None,
|
||||
):
|
||||
super().__init__()
|
||||
net = []
|
||||
if layers == 1:
|
||||
net.append(
|
||||
normed_activation_layer(
|
||||
in_features,
|
||||
out_features,
|
||||
use_norm=use_output_norm,
|
||||
activation=output_activation,
|
||||
device=device,
|
||||
)
|
||||
)
|
||||
else:
|
||||
if input_activation:
|
||||
net.append(get_activation(hidden_activation))
|
||||
net.append(
|
||||
normed_activation_layer(
|
||||
in_features,
|
||||
hidden_dim,
|
||||
use_norm=use_norm,
|
||||
activation=hidden_activation,
|
||||
device=device,
|
||||
)
|
||||
)
|
||||
for _ in range(layers - 2):
|
||||
net.append(
|
||||
normed_activation_layer(
|
||||
hidden_dim,
|
||||
hidden_dim,
|
||||
use_norm=use_norm,
|
||||
activation=hidden_activation,
|
||||
device=device,
|
||||
)
|
||||
)
|
||||
net.append(
|
||||
normed_activation_layer(
|
||||
hidden_dim,
|
||||
out_features,
|
||||
use_norm=use_output_norm,
|
||||
activation=output_activation,
|
||||
device=device,
|
||||
)
|
||||
)
|
||||
self.net = nn.Sequential(*net)
|
||||
|
||||
def forward(self, x):
|
||||
return self.net(x)
|
||||
|
||||
|
||||
class CriticNetwork(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
n_obs,
|
||||
n_act,
|
||||
hidden_dim=256,
|
||||
use_norm=True,
|
||||
use_encoder_norm=False,
|
||||
encoder_layers=1,
|
||||
head_layers=1,
|
||||
pred_layers=1,
|
||||
device=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.feature_module = FCNN(
|
||||
in_features=n_obs + n_act,
|
||||
out_features=hidden_dim,
|
||||
hidden_dim=hidden_dim,
|
||||
hidden_activation="swish",
|
||||
output_activation=None,
|
||||
use_norm=use_norm,
|
||||
use_output_norm=use_encoder_norm,
|
||||
layers=encoder_layers,
|
||||
device=device,
|
||||
)
|
||||
self.critic_module = FCNN(
|
||||
in_features=hidden_dim,
|
||||
out_features=1,
|
||||
hidden_dim=hidden_dim,
|
||||
hidden_activation="swish",
|
||||
output_activation=None,
|
||||
use_norm=use_norm,
|
||||
use_output_norm=False,
|
||||
layers=head_layers,
|
||||
device=device,
|
||||
)
|
||||
self.pred_module = FCNN(
|
||||
in_features=hidden_dim,
|
||||
out_features=hidden_dim,
|
||||
hidden_dim=hidden_dim,
|
||||
hidden_activation="swish",
|
||||
output_activation=None,
|
||||
use_norm=use_norm,
|
||||
use_output_norm=False,
|
||||
layers=pred_layers,
|
||||
device=device,
|
||||
)
|
||||
|
||||
def features(self, obs, action):
|
||||
state = torch.cat([obs, action], dim=-1)
|
||||
return self.feature_module(state)
|
||||
|
||||
def critic_head(self, features):
|
||||
return self.critic_module(features)
|
||||
|
||||
def critic(self, obs, action):
|
||||
features = self.features(obs, action)
|
||||
return self.critic_head(features)
|
||||
|
||||
def forward(self, obs, action):
|
||||
features = self.features(obs, action)
|
||||
return self.pred_module(features)
|
||||
|
||||
|
||||
class Critic(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
n_obs,
|
||||
n_act,
|
||||
num_atoms: int,
|
||||
vmin: float,
|
||||
vmax: float,
|
||||
hidden_dim=256,
|
||||
use_norm=True,
|
||||
use_encoder_norm=False,
|
||||
encoder_layers=1,
|
||||
head_layers=1,
|
||||
pred_layers=1,
|
||||
device=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_atoms = num_atoms
|
||||
self.vmin = vmin
|
||||
self.vmax = vmax
|
||||
self.hidden_dim = hidden_dim
|
||||
self.feature_module = FCNN(
|
||||
in_features=n_obs + n_act,
|
||||
out_features=hidden_dim,
|
||||
hidden_dim=hidden_dim,
|
||||
hidden_activation="swish",
|
||||
output_activation=None,
|
||||
use_norm=use_norm,
|
||||
use_output_norm=use_encoder_norm,
|
||||
layers=encoder_layers,
|
||||
device=device,
|
||||
)
|
||||
self.critic_module = FCNN(
|
||||
in_features=hidden_dim,
|
||||
out_features=num_atoms,
|
||||
hidden_dim=hidden_dim,
|
||||
hidden_activation="swish",
|
||||
output_activation=None,
|
||||
use_norm=use_norm,
|
||||
use_output_norm=False,
|
||||
input_activation=True,
|
||||
layers=head_layers,
|
||||
device=device,
|
||||
)
|
||||
self.pred_module = FCNN(
|
||||
in_features=hidden_dim,
|
||||
out_features=hidden_dim,
|
||||
hidden_dim=hidden_dim,
|
||||
hidden_activation="swish",
|
||||
output_activation=None,
|
||||
use_norm=use_norm,
|
||||
input_activation=True,
|
||||
use_output_norm=False,
|
||||
layers=pred_layers,
|
||||
device=device,
|
||||
)
|
||||
self.values = torch.linspace(
|
||||
vmin, vmax, num_atoms, device=device, dtype=torch.float32
|
||||
)
|
||||
zeros = hl_gauss(
|
||||
torch.zeros(1, device=device), self.vmin, self.vmax, self.num_atoms
|
||||
)
|
||||
zeros.requires_grad = True
|
||||
self.zero_dist = nn.Parameter(
|
||||
hl_gauss(
|
||||
torch.zeros(1, device=device), self.vmin, self.vmax, self.num_atoms
|
||||
)
|
||||
)
|
||||
|
||||
def forward(self, obs, action):
|
||||
inp = torch.cat([obs, action], dim=-1)
|
||||
features = self.feature_module(inp)
|
||||
next_pred = self.pred_module(features)
|
||||
logits = self.critic_module(features) + 40.9 * self.zero_dist
|
||||
value_cats = torch.softmax(logits, dim=-1)
|
||||
value = value_cats @ self.values
|
||||
return value, logits, next_pred, features
|
||||
|
||||
|
||||
class Actor(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
n_obs,
|
||||
n_act,
|
||||
ent_start: float,
|
||||
kl_start: float,
|
||||
hidden_dim=256,
|
||||
use_norm=True,
|
||||
layers=2,
|
||||
min_std=0.1,
|
||||
device=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.model = FCNN(
|
||||
in_features=n_obs,
|
||||
out_features=2 * n_act,
|
||||
hidden_dim=hidden_dim,
|
||||
hidden_activation="swish",
|
||||
output_activation=None,
|
||||
use_norm=use_norm,
|
||||
use_output_norm=False,
|
||||
layers=layers,
|
||||
device=device,
|
||||
)
|
||||
self.log_temp = nn.Parameter(
|
||||
torch.log(torch.tensor(ent_start, device=device, dtype=torch.float32))
|
||||
)
|
||||
self.log_lagrange = nn.Parameter(
|
||||
torch.log(torch.tensor(kl_start, device=device, dtype=torch.float32))
|
||||
)
|
||||
self.min_std = min_std
|
||||
|
||||
def forward(self, obs: torch.Tensor) -> torch.distributions.Distribution:
|
||||
x = self.model(obs)
|
||||
mean, log_std = torch.split(x, x.shape[-1] // 2, dim=-1)
|
||||
std = torch.exp(log_std) + self.min_std
|
||||
pi = Normal(mean, std, validate_args=False)
|
||||
|
||||
transformed_pi = torch.distributions.TransformedDistribution(
|
||||
pi, [torch.distributions.TanhTransform()]
|
||||
)
|
||||
return (
|
||||
transformed_pi,
|
||||
torch.tanh(mean),
|
||||
torch.exp(self.log_temp),
|
||||
torch.exp(self.log_lagrange),
|
||||
)
|
||||
|
||||
|
||||
class StochasticPolicy(nn.Module):
|
||||
def __init__(self, actor: Actor, normalizer: nn.Module = None, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.actor = actor
|
||||
self.normalizer = normalizer
|
||||
|
||||
def forward(self, obs: torch.Tensor) -> torch.distributions.Distribution:
|
||||
if self.normalizer:
|
||||
obs = self.normalizer(obs)
|
||||
return self.actor(obs)
|
||||
|
||||
|
||||
class TD3DeterministicPolicy(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
n_obs,
|
||||
n_act,
|
||||
hidden_dim=256,
|
||||
use_norm=True,
|
||||
layers=2,
|
||||
device=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.model = FCNN(
|
||||
in_features=n_obs,
|
||||
out_features=2 * n_act,
|
||||
hidden_dim=hidden_dim,
|
||||
hidden_activation="swish",
|
||||
output_activation=None,
|
||||
use_norm=use_norm,
|
||||
use_output_norm=False,
|
||||
layers=layers,
|
||||
device=device,
|
||||
)
|
||||
|
||||
def forward(self, obs: torch.Tensor) -> torch.Tensor:
|
||||
x = self.model(obs)
|
||||
mean, _ = torch.split(x, x.shape[-1] // 2, dim=-1)
|
||||
return torch.tanh(mean)
|
||||
Reference in New Issue
Block a user