Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
57 changes: 44 additions & 13 deletions ss2r/algorithms/mbpo/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -210,6 +210,7 @@ def train(
advantage_threshold: float = 0.2,
offline: bool = False,
learn_from_scratch: bool = False,
load_from_sbsrl: bool = False,
load_auxiliaries: bool = False,
load_normalizer: bool = True,
target_entropy: float | None = None,
Expand Down Expand Up @@ -374,6 +375,19 @@ def train(
backup_qc_params=params[4] if safe else None,
backup_target_qc_params=params[4] if safe else None,
)
elif load_from_sbsrl:
training_state = training_state.replace( # type: ignore
normalizer_params=ts_normalizer_params,
behavior_policy_params=params[1],
backup_policy_params=params[1],
behavior_qr_params=params[13],
behavior_target_qr_params=params[13],
backup_qr_params=params[13],
behavior_qc_params=params[14] if safe else None,
behavior_target_qc_params=params[14] if safe else None,
backup_qc_params=params[14] if safe else None,
backup_target_qc_params=params[14] if safe else None,
)
else:
training_state = training_state.replace( # type: ignore
normalizer_params=ts_normalizer_params,
Expand All @@ -397,21 +411,38 @@ def train(
alpha_optimizer_state = restore_state(
params[7], training_state.alpha_optimizer_state
)
qr_optimizer_state = restore_state(
params[8][1]["inner_state"]
if isinstance(params[8][1], dict)
else params[8],
training_state.behavior_qr_optimizer_state,
)
if not safe:
qc_optimizer_state = None
if not load_from_sbsrl:
qr_optimizer_state = restore_state(
params[8][1]["inner_state"]
if isinstance(params[8][1], dict)
else params[8],
training_state.behavior_qr_optimizer_state,
)
if not safe:
qc_optimizer_state = None
else:
qc_optimizer_state = restore_state(
params[9][1]["inner_state"]
if isinstance(params[9][1], dict)
else params[9],
training_state.backup_qc_optimizer_state,
)
else:
qc_optimizer_state = restore_state(
params[9][1]["inner_state"]
if isinstance(params[9][1], dict)
else params[9],
training_state.backup_qc_optimizer_state,
qr_optimizer_state = restore_state(
params[15][1]["inner_state"]
if isinstance(params[15][1], dict)
else params[15],
training_state.behavior_qr_optimizer_state,
)
if not safe:
qc_optimizer_state = None
else:
qc_optimizer_state = restore_state(
params[16][1]["inner_state"]
if isinstance(params[16][1], dict)
else params[16],
training_state.backup_qc_optimizer_state,
)
training_state = training_state.replace( # type: ignore
behavior_policy_optimizer_state=policy_optimizer_state,
alpha_optimizer_state=alpha_optimizer_state,
Expand Down
89 changes: 87 additions & 2 deletions ss2r/algorithms/sbsrl/losses.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@
See: https://arxiv.org/abs/1906.08253
"""

from typing import Any, TypeAlias
from typing import Any, Callable, Optional, TypeAlias

import jax
import jax.numpy as jnp
Expand All @@ -26,6 +26,7 @@
from brax.training.types import Params, PRNGKey

from ss2r.algorithms.penalizers import Penalizer
from ss2r.algorithms.sac.q_transforms import QTransformation as QTransformationSAC
from ss2r.algorithms.sbsrl.networks import SBSRLNetworks
from ss2r.algorithms.sbsrl.q_transforms import QTransformation

Expand All @@ -44,6 +45,7 @@ def make_losses(
normalize_fn,
ensemble_size,
safe,
save_sooper_backup,
use_mean_critic,
uncertainty_constraint,
uncertainty_epsilon,
Expand All @@ -56,6 +58,8 @@ def make_losses(
policy_network = sbsrl_network.policy_network
qr_network = sbsrl_network.qr_network
qc_network = sbsrl_network.qc_network
backup_qr_network = sbsrl_network.backup_qr_network
backup_qc_network = sbsrl_network.backup_qc_network
parametric_action_distribution = sbsrl_network.parametric_action_distribution

def alpha_loss(
Expand Down Expand Up @@ -317,4 +321,85 @@ def compute_model_loss(model_params, normalizer_params, data, obs_key="state"):
total_loss = jnp.mean(total_loss)
return total_loss

return alpha_loss, critic_loss_vmap, actor_loss, compute_model_loss
backup_critic_loss: Optional[
Callable[
[
Params,
Params,
Any,
Params,
jnp.ndarray,
Transition,
PRNGKey,
QTransformationSAC,
bool,
],
jnp.ndarray,
]
] = None
if backup_qc_network is not None and backup_qr_network is not None:

def backup_critic_loss(
q_params: Params,
policy_params: Params,
normalizer_params: Any,
target_q_params: Params,
alpha: jnp.ndarray,
transitions: Transition,
key: PRNGKey,
target_q_fn: QTransformationSAC,
safe: bool = False,
) -> jnp.ndarray:
q_network = (
backup_qc_network
if safe or uncertainty_constraint
else backup_qr_network
)
assert q_network is not None
action = transitions.action
scale = cost_scaling if safe else reward_scaling
gamma = safety_discounting if safe else discounting
q_old_action = q_network.apply(
normalizer_params, q_params, transitions.observation, action
)
key, another_key = jax.random.split(key)

def policy(obs: jax.Array) -> tuple[jax.Array, jax.Array]:
next_dist_params = policy_network.apply(
normalizer_params, policy_params, obs
)
next_action = parametric_action_distribution.sample_no_postprocessing(
next_dist_params, key
)
next_log_prob = parametric_action_distribution.log_prob(
next_dist_params, next_action
)
next_action = parametric_action_distribution.postprocess(next_action)
return next_action, next_log_prob

q_fn = lambda obs, action: q_network.apply(
normalizer_params, target_q_params, obs, action
)
target_q = target_q_fn(
transitions,
q_fn,
policy,
gamma,
alpha,
scale,
another_key,
)
q_error = q_old_action - jnp.expand_dims(target_q, -1)
# Better bootstrapping for truncated episodes.
truncation = transitions.extras["state_extras"]["truncation"]
q_error *= jnp.expand_dims(1 - truncation, -1)
q_loss = 0.5 * jnp.mean(jnp.square(q_error))
return q_loss

return (
alpha_loss,
critic_loss_vmap,
actor_loss,
compute_model_loss,
backup_critic_loss,
)
38 changes: 37 additions & 1 deletion ss2r/algorithms/sbsrl/networks.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@
from brax.training import distribution, networks, types
from flax import linen

from ss2r.algorithms.sac.networks import MLP, BroNet
from ss2r.algorithms.sac.networks import MLP, BroNet, make_q_network

ActivationFn = Callable[[jnp.ndarray], jnp.ndarray]
Initializer = Callable[..., Any]
Expand All @@ -43,6 +43,7 @@ def __call__(
n_critics: int = 2,
n_heads: int = 1,
safe: bool = False,
save_sooper_backup: bool = False,
uncertainty_constraint: bool = False,
use_bro: bool = True,
ensemble_size: int = 10,
Expand All @@ -56,6 +57,8 @@ class SBSRLNetworks:
policy_network: networks.FeedForwardNetwork
qr_network: networks.FeedForwardNetwork
qc_network: networks.FeedForwardNetwork | None
backup_qr_network: networks.FeedForwardNetwork | None
backup_qc_network: networks.FeedForwardNetwork | None
model_network: networks.FeedForwardNetwork
parametric_action_distribution: distribution.ParametricDistribution

Expand Down Expand Up @@ -207,6 +210,7 @@ def make_sbsrl_networks(
n_heads: int = 1,
safe: bool = False,
uncertainty_constraint: bool = False,
save_sooper_backup: bool = False,
ensemble_size: int = 10,
embedding_dim: int = 4,
) -> SBSRLNetworks:
Expand Down Expand Up @@ -256,6 +260,36 @@ def make_sbsrl_networks(
)
else:
qc_network = None
if save_sooper_backup:
backup_qc_network = make_q_network(
observation_size,
action_size,
preprocess_observations_fn=preprocess_observations_fn,
hidden_layer_sizes=value_hidden_layer_sizes,
activation=activation,
obs_key=value_obs_key,
use_bro=use_bro,
n_critics=n_critics,
n_heads=n_heads,
)
backup_old_apply = backup_qc_network.apply
backup_qc_network.apply = lambda *args, **kwargs: jnn.softplus(
backup_old_apply(*args, **kwargs)
)
backup_qr_network = make_q_network(
observation_size,
action_size,
preprocess_observations_fn=preprocess_observations_fn,
hidden_layer_sizes=value_hidden_layer_sizes,
activation=activation,
obs_key=value_obs_key,
use_bro=use_bro,
n_critics=n_critics,
n_heads=n_heads,
)
else:
backup_qc_network = None
backup_qr_network = None
model_network = make_world_model_ensemble(
observation_size,
action_size,
Expand All @@ -268,6 +302,8 @@ def make_sbsrl_networks(
policy_network=policy_network,
qr_network=qr_network,
qc_network=qc_network,
backup_qc_network=backup_qc_network,
backup_qr_network=backup_qr_network,
model_network=model_network,
parametric_action_distribution=parametric_action_distribution,
) # type: ignore
Loading