Fixes build errors due to name conflicts

This commit is contained in:
cvoelcker
2025-07-21 18:31:20 -04:00
parent 094ee0c5ba
commit e2f99648ae
26 changed files with 52 additions and 79 deletions
View File
+270
View File
@@ -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
low = torch.floor(b).long()
u = torch.ceil(b).long()
l_mask = torch.logical_and((u > 0), (low == u))
u_mask = torch.logical_and((low < (self.num_atoms - 1)), (low == u))
low = torch.where(l_mask, low - 1, low)
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, (low + offset).view(-1), (next_dist * (u.float() - b)).view(-1)
)
proj_dist.view(-1).index_add_(
0, (u + offset).view(-1), (next_dist * (b - low.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)
+424
View File
@@ -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_alg.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
+376
View File
@@ -0,0 +1,376 @@
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_alg.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)