diff --git a/ss2r/algorithms/mbpo/__init__.py b/ss2r/algorithms/mbpo/__init__.py index d94ac2d08..0133ff470 100644 --- a/ss2r/algorithms/mbpo/__init__.py +++ b/ss2r/algorithms/mbpo/__init__.py @@ -2,7 +2,10 @@ import ss2r.algorithms.mbpo.networks as mbpo_networks from ss2r.algorithms.sac.data import get_collection_fn -from ss2r.algorithms.sac.q_transforms import get_reward_q_transform +from ss2r.algorithms.sac.q_transforms import ( + get_cost_q_transform, + get_reward_q_transform, +) def get_train_fn(cfg, checkpoint_path, restore_checkpoint_path): @@ -54,6 +57,7 @@ def get_train_fn(cfg, checkpoint_path, restore_checkpoint_path): policy_obs_key=policy_obs_key, ) reward_q_transform = get_reward_q_transform(cfg) + cost_q_transform = get_cost_q_transform(cfg) data_collection = get_collection_fn(cfg) train_fn = functools.partial( mbpo.train, @@ -62,6 +66,7 @@ def get_train_fn(cfg, checkpoint_path, restore_checkpoint_path): network_factory=network_factory, checkpoint_logdir=checkpoint_path, reward_q_transform=reward_q_transform, + cost_q_transform=cost_q_transform, get_experience_fn=data_collection, restore_checkpoint_path=restore_checkpoint_path, ) diff --git a/ss2r/algorithms/mbpo/model_env.py b/ss2r/algorithms/mbpo/model_env.py index cf973d1ed..71e194c54 100644 --- a/ss2r/algorithms/mbpo/model_env.py +++ b/ss2r/algorithms/mbpo/model_env.py @@ -2,7 +2,6 @@ import jax.numpy as jnp from brax import envs from brax.envs import base -from brax.training.acme import running_statistics from ss2r.algorithms.sac.types import float32 @@ -13,17 +12,21 @@ class ModelBasedEnv(envs.Env): def __init__( self, transitions, - observation_size: int, - action_size: int, + observation_size, + action_size, model_network, model_params, - normalizer_params: running_statistics.RunningStatisticsState, - ensemble_selection: str = "mean", # "random", "mean", or "pessimistic" - safety_budget: float = float("inf"), + qc_network, + qc_params, + normalizer_params, + ensemble_selection="mean", # "random", "mean", or "pessimistic" + safety_budget=float("inf"), ): super().__init__() self.model_network = model_network self.model_params = model_params + self.qc_network = qc_network + self.qc_params = qc_params self.normalizer_params = normalizer_params self.ensemble_selection = ensemble_selection self.safety_budget = safety_budget @@ -79,11 +82,18 @@ def step(self, state: base.State, action: jax.Array) -> base.State: truncation = jnp.zeros_like(reward, dtype=jnp.float32) state.info["cost"] = cost state.info["truncation"] = truncation - if "cumulative_cost" in state.info: + if self.qc_network is not None: prev_cumulative_cost = state.info["cumulative_cost"] curr_discount = state.info.get("curr_discount", jnp.ones_like(reward)) accumulated_cost_for_transition = ( - prev_cumulative_cost + curr_discount * cost + prev_cumulative_cost + + curr_discount + * self.qc_network.apply( + self.normalizer_params, + self.qc_params, + state.obs, + action, + ).mean(axis=-1) ) done = jnp.where( accumulated_cost_for_transition > self.safety_budget, @@ -104,7 +114,6 @@ def reset_states(self, done, state, next_obs): return state, next_obs state, next_obs = reset_states(self, done, state, next_obs) - state = state.replace( obs=next_obs, reward=reward, @@ -130,11 +139,13 @@ def create_model_env( transitions, model_network, model_params, - observation_size: int, - action_size: int, - normalizer_params: running_statistics.RunningStatisticsState, - ensemble_selection: str = "random", - safety_budget: float = float("inf"), + qc_network, + qc_params, + observation_size, + action_size, + normalizer_params, + ensemble_selection="random", + safety_budget=float("inf"), ) -> ModelBasedEnv: """Factory function to create a model-based environment.""" return ModelBasedEnv( @@ -143,6 +154,8 @@ def create_model_env( observation_size=observation_size, action_size=action_size, model_params=model_params, + qc_network=qc_network, + qc_params=qc_params, normalizer_params=normalizer_params, ensemble_selection=ensemble_selection, safety_budget=safety_budget, diff --git a/ss2r/algorithms/mbpo/networks.py b/ss2r/algorithms/mbpo/networks.py index 6282b197c..174816895 100644 --- a/ss2r/algorithms/mbpo/networks.py +++ b/ss2r/algorithms/mbpo/networks.py @@ -19,6 +19,7 @@ import brax.training.agents.sac.networks as sac_networks import flax import jax +import jax.nn as jnn import jax.numpy as jnp from brax.training import distribution, networks, types from flax import linen @@ -41,6 +42,7 @@ def __call__( *, n_critics: int = 2, n_heads: int = 1, + safe: bool = False, use_bro: bool = True, ) -> NetworkType: pass @@ -50,6 +52,7 @@ def __call__( class MBPONetworks: policy_network: networks.FeedForwardNetwork qr_network: networks.FeedForwardNetwork + qc_network: networks.FeedForwardNetwork | None model_network: networks.FeedForwardNetwork parametric_action_distribution: distribution.ParametricDistribution @@ -126,6 +129,7 @@ def make_mbpo_networks( use_bro: bool = True, n_critics: int = 2, n_heads: int = 1, + safe: bool = False, ) -> MBPONetworks: """Make SAC networks.""" parametric_action_distribution = distribution.NormalTanhDistribution( @@ -150,6 +154,24 @@ def make_mbpo_networks( n_critics=n_critics, n_heads=n_heads, ) + if safe: + 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, + ) + old_apply = qc_network.apply + qc_network.apply = lambda *args, **kwargs: jnn.softplus( + old_apply(*args, **kwargs) + ) + else: + qc_network = None model_network = make_world_model_ensemble( observation_size, action_size, @@ -161,6 +183,7 @@ def make_mbpo_networks( return MBPONetworks( policy_network=policy_network, qr_network=qr_network, + qc_network=qc_network, model_network=model_network, parametric_action_distribution=parametric_action_distribution, ) # type: ignore diff --git a/ss2r/algorithms/mbpo/train.py b/ss2r/algorithms/mbpo/train.py index 90b497e1c..ac422852d 100644 --- a/ss2r/algorithms/mbpo/train.py +++ b/ss2r/algorithms/mbpo/train.py @@ -40,7 +40,7 @@ from ss2r.algorithms.ppo.wrappers import TrackOnlineCosts from ss2r.algorithms.sac import gradients from ss2r.algorithms.sac.data import collect_single_step -from ss2r.algorithms.sac.q_transforms import QTransformation, SACBase +from ss2r.algorithms.sac.q_transforms import QTransformation, SACBase, SACCost from ss2r.algorithms.sac.rae import RAEReplayBuffer from ss2r.algorithms.sac.types import ( CollectDataFn, @@ -59,6 +59,7 @@ def _init_training_state( alpha_optimizer: optax.GradientTransformation, policy_optimizer: optax.GradientTransformation, qr_optimizer: optax.GradientTransformation, + qc_optimizer: optax.GradientTransformation, model_optimizer: optax.GradientTransformation, model_ensemble_size: int, ) -> TrainingState: @@ -74,6 +75,13 @@ def _init_training_state( model_keys = jax.random.split(key_model, model_ensemble_size) model_params = init_model_ensemble(model_keys) model_optimizer_state = model_optimizer.init(model_params) + if mbpo_network.qc_network is not None: + qc_params = mbpo_network.qc_network.init(key_qr) + assert qc_optimizer is not None + qc_optimizer_state = qc_optimizer.init(qc_params) + else: + qc_params = None + qc_optimizer_state = None if isinstance(obs_size, Mapping): obs_shape = { k: specs.Array(v, jnp.dtype("float32")) for k, v in obs_size.items() @@ -86,7 +94,10 @@ def _init_training_state( policy_params=policy_params, qr_optimizer_state=qr_optimizer_state, qr_params=qr_params, + qc_optimizer_state=qc_optimizer_state, + qc_params=qc_params, target_qr_params=qr_params, + target_qc_params=qc_params, model_params=model_params, model_optimizer_state=model_optimizer_state, gradient_steps=jnp.zeros(()), @@ -146,12 +157,14 @@ def train( safe: bool = False, safety_budget: float = float("inf"), reward_q_transform: QTransformation = SACBase(), + cost_q_transform: QTransformation = SACCost(), use_bro: bool = True, normalize_budget: bool = True, reset_on_eval: bool = True, store_buffer: bool = False, use_rae: bool = False, optimism: float = 0.0, + pessimism: float = 0.0, model_propagation: str = "nominal", ): if min_replay_size >= num_timesteps: @@ -203,6 +216,7 @@ def train( observation_size=obs_size, action_size=action_size, preprocess_observations_fn=normalize_fn, + safe=safe, use_bro=use_bro, n_critics=n_critics, n_heads=n_heads, @@ -215,6 +229,7 @@ def train( ) policy_optimizer = make_optimizer(learning_rate, 1.0) qr_optimizer = make_optimizer(critic_learning_rate, 1.0) + qc_optimizer = make_optimizer(critic_learning_rate, 1.0) model_optimizer = make_optimizer(model_learning_rate, 1.0) if isinstance(obs_size, Mapping): dummy_obs = {k: jnp.zeros(v) for k, v in obs_size.items()} @@ -248,6 +263,7 @@ def train( alpha_optimizer=alpha_optimizer, policy_optimizer=policy_optimizer, qr_optimizer=qr_optimizer, + qc_optimizer=qc_optimizer, model_optimizer=model_optimizer, model_ensemble_size=model_ensemble_size, ) @@ -260,7 +276,8 @@ def train( training_state = training_state.replace( # type: ignore normalizer_params=params[0], policy_params=params[1], - qr_params=params[3], + qr_params=params[2], + qc_params=params[3], model_params=params[4], ) if len(params) >= 6 and use_rae: @@ -307,6 +324,12 @@ def train( critic_loss, qr_optimizer, pmap_axis_name=None ) ) + if safe: + cost_critic_update = gradients.gradient_update_fn( # pytype: disable=wrong-arg-types # jax-ndarray + critic_loss, qc_optimizer, pmap_axis_name=None + ) + else: + cost_critic_update = None model_update = ( gradients.gradient_update_fn( # pytype: disable=wrong-arg-types # jax-ndarray model_loss, model_optimizer, pmap_axis_name=None @@ -328,6 +351,7 @@ def train( make_model_env = functools.partial( create_model_env, model_network=mbpo_network.model_network, + qc_network=mbpo_network.qc_network, action_size=action_size, observation_size=obs_size, ensemble_selection=model_propagation, @@ -341,10 +365,13 @@ def train( sac_replay_buffer, alpha_update, critic_update, + cost_critic_update, model_update, actor_update, + safe, min_alpha, reward_q_transform, + cost_q_transform, model_grad_updates_per_step, critic_grad_updates_per_step, extra_fields, @@ -355,6 +382,7 @@ def train( unroll_length, num_model_rollouts, optimism, + pessimism, model_to_real_data_ratio, ) @@ -554,6 +582,7 @@ def training_epoch_with_timing( training_state.normalizer_params, training_state.policy_params, training_state.qr_params, + training_state.qc_params, training_state.model_params, ) if store_buffer: @@ -575,6 +604,7 @@ def training_epoch_with_timing( training_state.normalizer_params, training_state.policy_params, training_state.qr_params, + training_state.qc_params, training_state.model_params, ) if store_buffer: diff --git a/ss2r/algorithms/mbpo/training_step.py b/ss2r/algorithms/mbpo/training_step.py index 9b2e41e54..cd89ef34f 100644 --- a/ss2r/algorithms/mbpo/training_step.py +++ b/ss2r/algorithms/mbpo/training_step.py @@ -26,10 +26,13 @@ def make_training_step( sac_replay_buffer, alpha_update, critic_update, + cost_critic_update, model_update, actor_update, + safe, min_alpha, reward_q_transform, + cost_q_transform, model_grad_updates_per_step, critic_grad_updates_per_step, extra_fields, @@ -40,6 +43,7 @@ def make_training_step( unroll_length, num_model_rollouts, optimism, + pessimism, model_to_real_data_ratio, ): def critic_sgd_step( @@ -61,17 +65,46 @@ def critic_sgd_step( optimizer_state=training_state.qr_optimizer_state, params=training_state.qr_params, ) + if safe: + cost_critic_loss, qc_params, qc_optimizer_state = cost_critic_update( + training_state.qc_params, + training_state.policy_params, + training_state.normalizer_params, + training_state.target_qc_params, + alpha, + transitions, + key_critic, + cost_q_transform, + True, + optimizer_state=training_state.qc_optimizer_state, + params=training_state.qc_params, + ) + cost_metrics = { + "cost_critic_loss": cost_critic_loss, + } + else: + cost_metrics = {} + qc_params = None + qc_optimizer_state = None polyak = lambda target, new: jax.tree_util.tree_map( lambda x, y: x * (1 - tau) + y * tau, target, new ) new_target_qr_params = polyak(training_state.target_qr_params, qr_params) + if safe: + new_target_qc_params = polyak(training_state.target_qc_params, qc_params) + else: + new_target_qc_params = None metrics = { "critic_loss": critic_loss, + **cost_metrics, } new_training_state = training_state.replace( # type: ignore qr_optimizer_state=qr_optimizer_state, qr_params=qr_params, + qc_optimizer_state=qc_optimizer_state, + qc_params=qc_params, target_qr_params=new_target_qr_params, + target_qc_params=new_target_qc_params, gradient_steps=training_state.gradient_steps + 1, ) return (new_training_state, key), metrics @@ -199,6 +232,10 @@ def relabel_transitions( ) disagreement = next_obs_pred.std(axis=0).mean(-1) new_reward = reward.mean(0) + disagreement * optimism + if safe: + cost = cost.mean(0) + disagreement * pessimism + transitions.extras["state_extras"]["cost"] = cost + next_obs_pred = next_obs_pred.mean(0) return Transition( observation=transitions.observation, @@ -238,6 +275,7 @@ def training_step( ) planning_env = make_model_env( model_params=training_state.model_params, + qc_params=training_state.qc_params, normalizer_params=training_state.normalizer_params, transitions=transitions, ) diff --git a/ss2r/algorithms/mbpo/types.py b/ss2r/algorithms/mbpo/types.py index 3bb9864ec..babfbd172 100644 --- a/ss2r/algorithms/mbpo/types.py +++ b/ss2r/algorithms/mbpo/types.py @@ -13,9 +13,12 @@ class TrainingState: policy_params: Params qr_optimizer_state: optax.OptState qr_params: Params + qc_optimizer_state: optax.OptState | None + qc_params: Params | None model_params: Params model_optimizer_state: optax.OptState target_qr_params: Params + target_qc_params: Params | None gradient_steps: jnp.ndarray env_steps: jnp.ndarray alpha_optimizer_state: optax.OptState diff --git a/ss2r/algorithms/sac/q_transforms.py b/ss2r/algorithms/sac/q_transforms.py index f61fd83e3..22e41dd40 100644 --- a/ss2r/algorithms/sac/q_transforms.py +++ b/ss2r/algorithms/sac/q_transforms.py @@ -42,6 +42,30 @@ def __call__( return target_q +class PessimisticCostUpdate(QTransformation): + def __call__( + self, + transitions: Transition, + q_fn: Callable[[Params, jax.Array], jax.Array], + policy: Callable[[jax.Array], tuple[jax.Array, jax.Array]], + gamma: float, + alpha: jax.Array | None = None, + scale: float = 1.0, + key: jax.Array | None = None, + ): + next_action, _ = policy(transitions.next_observation) + next_q = q_fn(transitions.next_observation, next_action) + next_v = next_q.mean(axis=-1) + cost = transitions.extras["state_extras"]["cost"] + new_target_q = jax.lax.stop_gradient( + cost * scale + transitions.discount * gamma * next_v + ) + old_q = q_fn(transitions.observation, transitions.action).mean(axis=-1) + # TODO: check if works (intersection of models) + target_q = jax.lax.stop_gradient(jnp.minimum(new_target_q, old_q)) + return target_q + + class RAMU(QTransformation): """ https://arxiv.org/pdf/2301.12593 @@ -252,6 +276,8 @@ def get_cost_q_transform(cfg): robustness = RAMU(**cfg.agent.cost_robustness) elif cfg.agent.cost_robustness.name == "ucb_cost": robustness = UCBCost() + elif cfg.agent.cost_robustness.name == "pessimistic_cost_update": + robustness = PessimisticCostUpdate() else: raise ValueError("Unknown robustness") return robustness diff --git a/ss2r/configs/agent/cost_robustness/pessimistic_cost_update.yaml b/ss2r/configs/agent/cost_robustness/pessimistic_cost_update.yaml new file mode 100644 index 000000000..19147658b --- /dev/null +++ b/ss2r/configs/agent/cost_robustness/pessimistic_cost_update.yaml @@ -0,0 +1 @@ +name: pessimistic_cost_update \ No newline at end of file diff --git a/ss2r/configs/agent/mbpo.yaml b/ss2r/configs/agent/mbpo.yaml index 0025bf3a7..23b4ecd06 100644 --- a/ss2r/configs/agent/mbpo.yaml +++ b/ss2r/configs/agent/mbpo.yaml @@ -1,5 +1,5 @@ defaults: - - cost_robustness: null + - cost_robustness: pess_cost_update - reward_robustness: null - propagation: null - data_collection: step @@ -41,4 +41,5 @@ model_ensemble_size: 5 unroll_length: 1 num_model_rollouts: 400 optimism: 0. +pessimism: 0. model_propagation: random diff --git a/ss2r/configs/experiment/cartpole_mbpo.yaml b/ss2r/configs/experiment/cartpole_mbpo.yaml index 40f8e735c..4989fbbe7 100644 --- a/ss2r/configs/experiment/cartpole_mbpo.yaml +++ b/ss2r/configs/experiment/cartpole_mbpo.yaml @@ -6,9 +6,9 @@ defaults: training: num_timesteps: 150000 - safe: true - safety_budget: 100 num_envs: 10 + safe: false + safety_budget: 100 train_domain_randomization: false eval_domain_randomization: false