From 90ecde3661e97c247c1a5a08258536985348de9c Mon Sep 17 00:00:00 2001 From: ManuelWendl Date: Mon, 9 Jun 2025 15:53:04 +0200 Subject: [PATCH 1/8] accumulated cost and termination --- ss2r/algorithms/mbpo/model_env.py | 11 ++++------- ss2r/algorithms/mbpo/train.py | 11 ++++++++++- ss2r/algorithms/mbpo/training_step.py | 17 +++++++++-------- ss2r/algorithms/ppo/wrappers.py | 9 ++++++++- ss2r/configs/experiment/cartpole_mbpo.yaml | 3 ++- 5 files changed, 33 insertions(+), 18 deletions(-) diff --git a/ss2r/algorithms/mbpo/model_env.py b/ss2r/algorithms/mbpo/model_env.py index d1b56a8df..5f377582d 100644 --- a/ss2r/algorithms/mbpo/model_env.py +++ b/ss2r/algorithms/mbpo/model_env.py @@ -54,19 +54,16 @@ def step(self, state: base.State, action: jax.Array) -> base.State: state.info["truncation"] = truncation if "cumulative_cost" in state.info: prev_cumulative_cost = state.info["cumulative_cost"] - accumulated_cost_for_transition = prev_cumulative_cost + cost + curr_discount = state.info.get("curr_discount", jnp.ones_like(reward)) + accumulated_cost_for_transition = ( + prev_cumulative_cost + curr_discount * cost + ) if self.safety_budget < float("inf"): done = jnp.where( accumulated_cost_for_transition > self.safety_budget, jnp.ones_like(done), done, ) - accumulated_cost_for_transition = jnp.where( - done > 0, - jnp.zeros_like(accumulated_cost_for_transition), - accumulated_cost_for_transition, - ) - state.info["cumulative_cost"] = accumulated_cost_for_transition state = state.replace( obs=next_obs, reward=reward, diff --git a/ss2r/algorithms/mbpo/train.py b/ss2r/algorithms/mbpo/train.py index 68f8cb900..90b497e1c 100644 --- a/ss2r/algorithms/mbpo/train.py +++ b/ss2r/algorithms/mbpo/train.py @@ -37,6 +37,7 @@ from ss2r.algorithms.mbpo.model_env import create_model_env from ss2r.algorithms.mbpo.training_step import make_training_step from ss2r.algorithms.mbpo.types import TrainingState +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 @@ -190,6 +191,8 @@ def train( env = environment if wrap_env_fn is not None: env = wrap_env_fn(env) + if safe: + env = TrackOnlineCosts(env, safety_discounting) rng = jax.random.PRNGKey(seed) obs_size = env.observation_size action_size = env.action_size @@ -226,6 +229,8 @@ def train( } if safe: extras["state_extras"]["cost"] = jnp.zeros(()) # type: ignore + extras["state_extras"]["cumulative_cost"] = jnp.zeros(()) # type: ignore + extras["state_extras"]["curr_discount"] = jnp.ones(()) # type: ignore dummy_transition = Transition( # pytype: disable=wrong-arg-types # jax-ndarray observation=dummy_obs, action=dummy_action, @@ -314,7 +319,11 @@ def train( ) extra_fields = ("truncation",) if safe: - extra_fields += ("cost",) # type: ignore + extra_fields += ( + "cost", + "cumulative_cost", + "curr_discount", + ) # type: ignore make_model_env = functools.partial( create_model_env, diff --git a/ss2r/algorithms/mbpo/training_step.py b/ss2r/algorithms/mbpo/training_step.py index 165f105d1..31a96ab60 100644 --- a/ss2r/algorithms/mbpo/training_step.py +++ b/ss2r/algorithms/mbpo/training_step.py @@ -174,20 +174,21 @@ def generate_model_data( ), "num_model_rollouts must be less than or equal to the number of transitions" transitions = jax.tree_map(lambda x: x[:num_model_rollouts], transitions) transitions = float32(transitions) - # FIXME: not zeros, use what happened in the real system - cumulative_cost = jax.random.uniform( - cost_key, (transitions.reward.shape[0],), minval=0.0, maxval=0.0 - ) state = envs.State( pipeline_state=None, obs=transitions.observation, reward=transitions.reward, - done=jnp.zeros_like(transitions.reward), + done=1 - transitions.discount, info={ - "cumulative_cost": cumulative_cost, # type: ignore - "truncation": jnp.zeros_like(cumulative_cost), + "cumulative_cost": transitions.extras["state_extras"].get( + "cumulative_cost", jnp.zeros_like(transitions.reward) + ), + "curr_discount": transitions.extras["state_extras"].get( + "curr_discount", jnp.ones_like(transitions.reward) + ), + "truncation": jnp.zeros_like(transitions.reward), "cost": transitions.extras["state_extras"].get( - "cost", jnp.zeros_like(cumulative_cost) + "cost", jnp.zeros_like(transitions.reward) ), "key": jnp.tile(model_key[None], (transitions.observation.shape[0], 1)), }, diff --git a/ss2r/algorithms/ppo/wrappers.py b/ss2r/algorithms/ppo/wrappers.py index 9c45bf849..b42e43a4c 100644 --- a/ss2r/algorithms/ppo/wrappers.py +++ b/ss2r/algorithms/ppo/wrappers.py @@ -4,11 +4,16 @@ class TrackOnlineCosts(Wrapper): + def __init__(self, env, cost_discount=1.0): + super().__init__(env) + self.cost_discount = cost_discount + def reset(self, rng: jax.Array) -> State: reset_state = self.env.reset(rng) reset_state.info["cumulative_cost"] = reset_state.info.get( "cost", jnp.zeros_like(reset_state.reward) ) + reset_state.info["curr_discount"] = jnp.ones_like(reset_state.reward) return reset_state def step(self, state: State, action: jax.Array) -> State: @@ -19,7 +24,9 @@ def step(self, state: State, action: jax.Array) -> State: ) nstate = self.env.step(state, action) cost = nstate.info.get("cost", jnp.zeros_like(nstate.reward)) - nstate.info.update(cumulative_cost=cumulative_cost + cost) + curr_discount = nstate.info.get("curr_discount", jnp.ones_like(nstate.reward)) + nstate.info.update(cumulative_cost=cumulative_cost + curr_discount * cost) + nstate.info.update(curr_discount=curr_discount * self.cost_discount) return nstate diff --git a/ss2r/configs/experiment/cartpole_mbpo.yaml b/ss2r/configs/experiment/cartpole_mbpo.yaml index a9b472ddb..40f8e735c 100644 --- a/ss2r/configs/experiment/cartpole_mbpo.yaml +++ b/ss2r/configs/experiment/cartpole_mbpo.yaml @@ -6,7 +6,8 @@ defaults: training: num_timesteps: 150000 - safe: false + safe: true + safety_budget: 100 num_envs: 10 train_domain_randomization: false eval_domain_randomization: false From d3b917a832710a2c1a91e926d9722c77215e31b6 Mon Sep 17 00:00:00 2001 From: ManuelWendl Date: Mon, 9 Jun 2025 16:17:40 +0200 Subject: [PATCH 2/8] test wo termination --- ss2r/algorithms/mbpo/model_env.py | 24 ++++++++++++------------ 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/ss2r/algorithms/mbpo/model_env.py b/ss2r/algorithms/mbpo/model_env.py index 5f377582d..245025c6f 100644 --- a/ss2r/algorithms/mbpo/model_env.py +++ b/ss2r/algorithms/mbpo/model_env.py @@ -52,18 +52,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: - 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 - ) - if self.safety_budget < float("inf"): - done = jnp.where( - accumulated_cost_for_transition > self.safety_budget, - jnp.ones_like(done), - done, - ) + # if "cumulative_cost" in state.info: + # 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 + # ) + # if self.safety_budget < float("inf"): + # done = jnp.where( + # accumulated_cost_for_transition > self.safety_budget, + # jnp.ones_like(done), + # done, + # ) state = state.replace( obs=next_obs, reward=reward, From be48f9da33dc0b2e2fb9100f7174876b287f2788 Mon Sep 17 00:00:00 2001 From: ManuelWendl Date: Mon, 9 Jun 2025 16:29:20 +0200 Subject: [PATCH 3/8] with termination --- ss2r/algorithms/mbpo/model_env.py | 24 ++++++++++++------------ 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/ss2r/algorithms/mbpo/model_env.py b/ss2r/algorithms/mbpo/model_env.py index 245025c6f..5f377582d 100644 --- a/ss2r/algorithms/mbpo/model_env.py +++ b/ss2r/algorithms/mbpo/model_env.py @@ -52,18 +52,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: - # 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 - # ) - # if self.safety_budget < float("inf"): - # done = jnp.where( - # accumulated_cost_for_transition > self.safety_budget, - # jnp.ones_like(done), - # done, - # ) + if "cumulative_cost" in state.info: + 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 + ) + if self.safety_budget < float("inf"): + done = jnp.where( + accumulated_cost_for_transition > self.safety_budget, + jnp.ones_like(done), + done, + ) state = state.replace( obs=next_obs, reward=reward, From c9552ad904e39ec0fdad7f8a9b6bafe92792eb2c Mon Sep 17 00:00:00 2001 From: ManuelWendl Date: Mon, 9 Jun 2025 18:32:20 +0200 Subject: [PATCH 4/8] reset of model_env --- ss2r/algorithms/mbpo/model_env.py | 57 +++++++++++++++++++++++---- ss2r/algorithms/mbpo/training_step.py | 37 +++++------------ 2 files changed, 60 insertions(+), 34 deletions(-) diff --git a/ss2r/algorithms/mbpo/model_env.py b/ss2r/algorithms/mbpo/model_env.py index 5f377582d..cf973d1ed 100644 --- a/ss2r/algorithms/mbpo/model_env.py +++ b/ss2r/algorithms/mbpo/model_env.py @@ -4,12 +4,15 @@ from brax.envs import base from brax.training.acme import running_statistics +from ss2r.algorithms.sac.types import float32 + class ModelBasedEnv(envs.Env): """Environment wrapper that uses a learned model for predictions.""" def __init__( self, + transitions, observation_size: int, action_size: int, model_network, @@ -26,12 +29,36 @@ def __init__( self.safety_budget = safety_budget self._observation_size = observation_size self._action_size = action_size + self.transitions = transitions def reset(self, rng: jax.Array) -> base.State: - """Reset using the real environment.""" - raise NotImplementedError( - "ModelBasedEnv does not support reset. Use a real environment for resetting." + sample_key, model_key = jax.random.split(rng) + indcs = jax.random.randint( + sample_key, (), 0, self.transitions.observation.shape[0] + ) + transitions = float32( + jax.tree_util.tree_map(lambda x: x[indcs], self.transitions) ) + state = envs.State( + pipeline_state=None, + obs=transitions.observation, + reward=transitions.reward, + done=jnp.zeros_like(transitions.reward), + info={ + "cumulative_cost": transitions.extras["state_extras"].get( + "cumulative_cost", jnp.zeros_like(transitions.reward) + ), + "curr_discount": transitions.extras["state_extras"].get( + "curr_discount", jnp.ones_like(transitions.reward) + ), + "truncation": jnp.zeros_like(transitions.reward), + "cost": transitions.extras["state_extras"].get( + "cost", jnp.zeros_like(transitions.reward) + ), + "key": model_key, + }, + ) + return state def step(self, state: base.State, action: jax.Array) -> base.State: """Step using the learned model.""" @@ -58,12 +85,26 @@ def step(self, state: base.State, action: jax.Array) -> base.State: accumulated_cost_for_transition = ( prev_cumulative_cost + curr_discount * cost ) - if self.safety_budget < float("inf"): - done = jnp.where( - accumulated_cost_for_transition > self.safety_budget, - jnp.ones_like(done), + done = jnp.where( + accumulated_cost_for_transition > self.safety_budget, + jnp.ones_like(done), + done, + ) + + def reset_states(self, done, state, next_obs): + """Reset the state if done.""" + key, reset_keys = jax.random.split(state.info["key"]) + state.info["key"] = key + state.info["cumulative_cost"] = jnp.where( done, + jnp.zeros_like(reward), + state.info["cumulative_cost"], ) + next_obs = jnp.where(done, self.reset(reset_keys).obs, next_obs) + return state, next_obs + + state, next_obs = reset_states(self, done, state, next_obs) + state = state.replace( obs=next_obs, reward=reward, @@ -86,6 +127,7 @@ def backend(self) -> str: def create_model_env( + transitions, model_network, model_params, observation_size: int, @@ -96,6 +138,7 @@ def create_model_env( ) -> ModelBasedEnv: """Factory function to create a model-based environment.""" return ModelBasedEnv( + transitions=transitions, model_network=model_network, observation_size=observation_size, action_size=action_size, diff --git a/ss2r/algorithms/mbpo/training_step.py b/ss2r/algorithms/mbpo/training_step.py index 31a96ab60..9b2e41e54 100644 --- a/ss2r/algorithms/mbpo/training_step.py +++ b/ss2r/algorithms/mbpo/training_step.py @@ -163,36 +163,15 @@ def run_experience_step( def generate_model_data( planning_env: ModelBasedEnv, policy: Policy, - transitions: Transition, sac_replay_buffer_state: ReplayBufferState, key: PRNGKey, ) -> ReplayBufferState: key_generate_unroll, cost_key, model_key, key_perm = jax.random.split(key, 4) - assert ( - num_model_rollouts - <= transitions.observation.shape[0] * transitions.observation.shape[1] - ), "num_model_rollouts must be less than or equal to the number of transitions" - transitions = jax.tree_map(lambda x: x[:num_model_rollouts], transitions) - transitions = float32(transitions) - state = envs.State( - pipeline_state=None, - obs=transitions.observation, - reward=transitions.reward, - done=1 - transitions.discount, - info={ - "cumulative_cost": transitions.extras["state_extras"].get( - "cumulative_cost", jnp.zeros_like(transitions.reward) - ), - "curr_discount": transitions.extras["state_extras"].get( - "curr_discount", jnp.ones_like(transitions.reward) - ), - "truncation": jnp.zeros_like(transitions.reward), - "cost": transitions.extras["state_extras"].get( - "cost", jnp.zeros_like(transitions.reward) - ), - "key": jnp.tile(model_key[None], (transitions.observation.shape[0], 1)), - }, - ) + keys = jax.random.split(key, num_model_rollouts + 2) + key = keys[0] + key_generate_unroll = keys[1] + rollout_keys = keys[2:] + state = planning_env.reset(rollout_keys) _, transitions = acting.generate_unroll( planning_env, state, @@ -245,6 +224,9 @@ def training_step( training_key, ) = run_experience_step(training_state, env_state, model_buffer_state, key) model_buffer_state, transitions = model_replay_buffer.sample(model_buffer_state) + assert ( + num_model_rollouts <= transitions.observation.shape[0] + ), "num_model_rollouts must be less than or equal to the number of transitions" # Change the front dimension of transitions so 'update_step' is called # grad_updates_per_step times by the scan. tmp_transitions = jax.tree_util.tree_map( @@ -257,6 +239,7 @@ def training_step( planning_env = make_model_env( model_params=training_state.model_params, normalizer_params=training_state.normalizer_params, + transitions=transitions, ) planning_env = VmapWrapper(planning_env) policy = make_policy( @@ -264,7 +247,7 @@ def training_step( ) # Rollout trajectories from the sampled transitions sac_buffer_state = generate_model_data( - planning_env, policy, transitions, sac_buffer_state, training_key + planning_env, policy, sac_buffer_state, training_key ) # Train SAC with model data sac_buffer_state, model_transitions = sac_replay_buffer.sample(sac_buffer_state) From b06ec90fad0a5387a2129c40945fd516f9b3391d Mon Sep 17 00:00:00 2001 From: ManuelWendl Date: Tue, 10 Jun 2025 08:43:02 +0200 Subject: [PATCH 5/8] cost-critic: pessimistic update using disagreement for cost, planning mdp termination with q function --- ss2r/algorithms/mbpo/__init__.py | 7 ++- ss2r/algorithms/mbpo/model_env.py | 17 +++++- ss2r/algorithms/mbpo/networks.py | 13 +++++ ss2r/algorithms/mbpo/train.py | 28 +++++++++- ss2r/algorithms/mbpo/training_step.py | 65 +++++++++++++++++++++- ss2r/algorithms/mbpo/types.py | 3 + ss2r/algorithms/sac/q_transforms.py | 29 ++++++++++ ss2r/configs/agent/mbpo.yaml | 3 +- ss2r/configs/experiment/cartpole_mbpo.yaml | 2 - 9 files changed, 159 insertions(+), 8 deletions(-) diff --git a/ss2r/algorithms/mbpo/__init__.py b/ss2r/algorithms/mbpo/__init__.py index 9bbef039f..296044780 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): @@ -53,6 +56,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, @@ -61,6 +65,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..33cb55b7e 100644 --- a/ss2r/algorithms/mbpo/model_env.py +++ b/ss2r/algorithms/mbpo/model_env.py @@ -17,6 +17,8 @@ def __init__( action_size: int, model_network, model_params, + qc_network, + qc_params, normalizer_params: running_statistics.RunningStatisticsState, ensemble_selection: str = "mean", # "random", "mean", or "pessimistic" safety_budget: float = float("inf"), @@ -24,6 +26,8 @@ def __init__( 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 @@ -83,7 +87,14 @@ def step(self, state: base.State, action: jax.Array) -> base.State: 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, + ).max(axis=-1) ) done = jnp.where( accumulated_cost_for_transition > self.safety_budget, @@ -130,6 +141,8 @@ def create_model_env( transitions, model_network, model_params, + qc_network, + qc_params, observation_size: int, action_size: int, normalizer_params: running_statistics.RunningStatisticsState, @@ -143,6 +156,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..a87a70715 100644 --- a/ss2r/algorithms/mbpo/networks.py +++ b/ss2r/algorithms/mbpo/networks.py @@ -50,6 +50,7 @@ def __call__( class MBPONetworks: policy_network: networks.FeedForwardNetwork qr_network: networks.FeedForwardNetwork + qc_network: networks.FeedForwardNetwork model_network: networks.FeedForwardNetwork parametric_action_distribution: distribution.ParametricDistribution @@ -150,6 +151,17 @@ def make_mbpo_networks( n_critics=n_critics, n_heads=n_heads, ) + 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, + ) model_network = make_world_model_ensemble( observation_size, action_size, @@ -161,6 +173,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..8dd3a09eb 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: @@ -70,6 +71,8 @@ def _init_training_state( policy_optimizer_state = policy_optimizer.init(policy_params) qr_params = mbpo_network.qr_network.init(key_qr) qr_optimizer_state = qr_optimizer.init(qr_params) + qc_params = mbpo_network.qc_network.init(key_qr) + qc_optimizer_state = qc_optimizer.init(qc_params) init_model_ensemble = jax.vmap(mbpo_network.model_network.init) model_keys = jax.random.split(key_model, model_ensemble_size) model_params = init_model_ensemble(model_keys) @@ -86,7 +89,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 +152,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: @@ -215,6 +223,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 +257,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 +270,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 +318,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 +345,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 +359,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 +376,7 @@ def train( unroll_length, num_model_rollouts, optimism, + pessimism, model_to_real_data_ratio, ) @@ -554,6 +576,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 +598,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..cd58ae784 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,13 +232,42 @@ 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 + # TODO: Check if works (Reduce conservatism of terminations) + acc_cost = transitions.extras["state_extras"].get("accumulated_cost", 0.0) + future_cost = planning_env.qc_network.apply( + normalizer_params, + planning_env.qc_params, + transitions.observation, + transitions.action, + ).mean(axis=-1) + new_discount = jnp.where( # Inverted done signal + acc_cost + + transitions.extras["state_extras"].get( + "curr_discount", jnp.ones_like(transitions.reward) + ) + * future_cost + > planning_env.safety_budget, + jnp.zeros_like(cost, dtype=jnp.float32), + jnp.ones_like(cost, dtype=jnp.float32), + ) + discount = ( + jnp.logical_or( # Logical OR enforce more ones -> reduce pessimism + transitions.discount, new_discount + ) + ) + else: + discount = transitions.discount + next_obs_pred = next_obs_pred.mean(0) return Transition( observation=transitions.observation, next_observation=next_obs_pred, action=transitions.action, reward=new_reward, - discount=transitions.discount, + discount=discount, extras=transitions.extras, ), disagreement @@ -238,6 +300,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..809431ad7 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 + qc_params: Params model_params: Params model_optimizer_state: optax.OptState target_qr_params: Params + target_qc_params: Params 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..53655b85d 100644 --- a/ss2r/algorithms/sac/q_transforms.py +++ b/ss2r/algorithms/sac/q_transforms.py @@ -42,6 +42,33 @@ def __call__( return target_q +class PessCostUpdate(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) + target_q = jax.lax.stop_gradient( + jnp.minimum( + new_target_q, old_q + ) # TODO: check if works (intersection of models) + ) + return target_q + + class RAMU(QTransformation): """ https://arxiv.org/pdf/2301.12593 @@ -252,6 +279,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 == "pess_cost_update": + robustness = PessCostUpdate() else: raise ValueError("Unknown robustness") return robustness 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..ab4a618eb 100644 --- a/ss2r/configs/experiment/cartpole_mbpo.yaml +++ b/ss2r/configs/experiment/cartpole_mbpo.yaml @@ -6,8 +6,6 @@ defaults: training: num_timesteps: 150000 - safe: true - safety_budget: 100 num_envs: 10 train_domain_randomization: false eval_domain_randomization: false From 259c50966aaa08cbb8fd0c88f6cfa1520837c2b1 Mon Sep 17 00:00:00 2001 From: ManuelWendl Date: Tue, 10 Jun 2025 08:57:37 +0200 Subject: [PATCH 6/8] configs --- .../cost_robustness/pess_cost_update.yaml | 1 + .../experiment/cartpole_mbpo_safe.yaml | 21 +++++++++++++++++++ 2 files changed, 22 insertions(+) create mode 100644 ss2r/configs/agent/cost_robustness/pess_cost_update.yaml create mode 100644 ss2r/configs/experiment/cartpole_mbpo_safe.yaml diff --git a/ss2r/configs/agent/cost_robustness/pess_cost_update.yaml b/ss2r/configs/agent/cost_robustness/pess_cost_update.yaml new file mode 100644 index 000000000..f614a125c --- /dev/null +++ b/ss2r/configs/agent/cost_robustness/pess_cost_update.yaml @@ -0,0 +1 @@ +name: pess_cost_update \ No newline at end of file diff --git a/ss2r/configs/experiment/cartpole_mbpo_safe.yaml b/ss2r/configs/experiment/cartpole_mbpo_safe.yaml new file mode 100644 index 000000000..22cbe0ecd --- /dev/null +++ b/ss2r/configs/experiment/cartpole_mbpo_safe.yaml @@ -0,0 +1,21 @@ +# @package _global_ +defaults: + - override /environment: cartpole + - override /agent: mbpo + - override /agent/cost_robustness: pess_cost_update + - _self_ + +training: + num_timesteps: 150000 + safe: true + safety_budget: 100 + num_envs: 10 + train_domain_randomization: false + eval_domain_randomization: false + +agent: + min_replay_size: 5000 + sac_batch_size: 512 + critic_grad_updates_per_step: 20 + model_grad_updates_per_step: 25 + pessimism: 0.5 \ No newline at end of file From 38026d0111349d4e8da754bef100f2261a9a0031 Mon Sep 17 00:00:00 2001 From: ManuelWendl Date: Tue, 10 Jun 2025 09:45:09 +0200 Subject: [PATCH 7/8] fix safe=false case, None types consistent --- ss2r/algorithms/mbpo/model_env.py | 2 +- ss2r/algorithms/mbpo/networks.py | 32 ++++++++++++------- ss2r/algorithms/mbpo/train.py | 10 ++++-- .../experiment/cartpole_mbpo_safe.yaml | 2 +- 4 files changed, 31 insertions(+), 15 deletions(-) diff --git a/ss2r/algorithms/mbpo/model_env.py b/ss2r/algorithms/mbpo/model_env.py index 33cb55b7e..01adf498c 100644 --- a/ss2r/algorithms/mbpo/model_env.py +++ b/ss2r/algorithms/mbpo/model_env.py @@ -83,7 +83,7 @@ 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 = ( diff --git a/ss2r/algorithms/mbpo/networks.py b/ss2r/algorithms/mbpo/networks.py index a87a70715..c954e30f3 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 @@ -127,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( @@ -151,17 +154,24 @@ def make_mbpo_networks( n_critics=n_critics, n_heads=n_heads, ) - 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, - ) + 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, diff --git a/ss2r/algorithms/mbpo/train.py b/ss2r/algorithms/mbpo/train.py index 8dd3a09eb..ac422852d 100644 --- a/ss2r/algorithms/mbpo/train.py +++ b/ss2r/algorithms/mbpo/train.py @@ -71,12 +71,17 @@ def _init_training_state( policy_optimizer_state = policy_optimizer.init(policy_params) qr_params = mbpo_network.qr_network.init(key_qr) qr_optimizer_state = qr_optimizer.init(qr_params) - qc_params = mbpo_network.qc_network.init(key_qr) - qc_optimizer_state = qc_optimizer.init(qc_params) init_model_ensemble = jax.vmap(mbpo_network.model_network.init) 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() @@ -211,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, diff --git a/ss2r/configs/experiment/cartpole_mbpo_safe.yaml b/ss2r/configs/experiment/cartpole_mbpo_safe.yaml index 22cbe0ecd..5d459d22b 100644 --- a/ss2r/configs/experiment/cartpole_mbpo_safe.yaml +++ b/ss2r/configs/experiment/cartpole_mbpo_safe.yaml @@ -7,7 +7,7 @@ defaults: training: num_timesteps: 150000 - safe: true + safe: false safety_budget: 100 num_envs: 10 train_domain_randomization: false From 28473f503ef74a44bb3fea3c21da0a9331afde48 Mon Sep 17 00:00:00 2001 From: Yarden Date: Tue, 10 Jun 2025 12:40:45 +0200 Subject: [PATCH 8/8] Updates after review --- ss2r/algorithms/mbpo/model_env.py | 24 ++++++++--------- ss2r/algorithms/mbpo/networks.py | 2 +- ss2r/algorithms/mbpo/training_step.py | 27 +------------------ ss2r/algorithms/mbpo/types.py | 6 ++--- ss2r/algorithms/sac/q_transforms.py | 13 ++++----- .../cost_robustness/pess_cost_update.yaml | 1 - .../pessimistic_cost_update.yaml | 1 + ss2r/configs/experiment/cartpole_mbpo.yaml | 2 ++ .../experiment/cartpole_mbpo_safe.yaml | 21 --------------- 9 files changed, 24 insertions(+), 73 deletions(-) delete mode 100644 ss2r/configs/agent/cost_robustness/pess_cost_update.yaml create mode 100644 ss2r/configs/agent/cost_robustness/pessimistic_cost_update.yaml delete mode 100644 ss2r/configs/experiment/cartpole_mbpo_safe.yaml diff --git a/ss2r/algorithms/mbpo/model_env.py b/ss2r/algorithms/mbpo/model_env.py index 01adf498c..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,15 +12,15 @@ class ModelBasedEnv(envs.Env): def __init__( self, transitions, - observation_size: int, - action_size: int, + observation_size, + action_size, model_network, model_params, qc_network, qc_params, - normalizer_params: running_statistics.RunningStatisticsState, - ensemble_selection: str = "mean", # "random", "mean", or "pessimistic" - safety_budget: float = float("inf"), + normalizer_params, + ensemble_selection="mean", # "random", "mean", or "pessimistic" + safety_budget=float("inf"), ): super().__init__() self.model_network = model_network @@ -94,7 +93,7 @@ def step(self, state: base.State, action: jax.Array) -> base.State: self.qc_params, state.obs, action, - ).max(axis=-1) + ).mean(axis=-1) ) done = jnp.where( accumulated_cost_for_transition > self.safety_budget, @@ -115,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, @@ -143,11 +141,11 @@ def create_model_env( model_params, qc_network, qc_params, - observation_size: int, - action_size: int, - normalizer_params: running_statistics.RunningStatisticsState, - ensemble_selection: str = "random", - safety_budget: float = float("inf"), + observation_size, + action_size, + normalizer_params, + ensemble_selection="random", + safety_budget=float("inf"), ) -> ModelBasedEnv: """Factory function to create a model-based environment.""" return ModelBasedEnv( diff --git a/ss2r/algorithms/mbpo/networks.py b/ss2r/algorithms/mbpo/networks.py index c954e30f3..174816895 100644 --- a/ss2r/algorithms/mbpo/networks.py +++ b/ss2r/algorithms/mbpo/networks.py @@ -52,7 +52,7 @@ def __call__( class MBPONetworks: policy_network: networks.FeedForwardNetwork qr_network: networks.FeedForwardNetwork - qc_network: networks.FeedForwardNetwork + qc_network: networks.FeedForwardNetwork | None model_network: networks.FeedForwardNetwork parametric_action_distribution: distribution.ParametricDistribution diff --git a/ss2r/algorithms/mbpo/training_step.py b/ss2r/algorithms/mbpo/training_step.py index cd58ae784..cd89ef34f 100644 --- a/ss2r/algorithms/mbpo/training_step.py +++ b/ss2r/algorithms/mbpo/training_step.py @@ -235,31 +235,6 @@ def relabel_transitions( if safe: cost = cost.mean(0) + disagreement * pessimism transitions.extras["state_extras"]["cost"] = cost - # TODO: Check if works (Reduce conservatism of terminations) - acc_cost = transitions.extras["state_extras"].get("accumulated_cost", 0.0) - future_cost = planning_env.qc_network.apply( - normalizer_params, - planning_env.qc_params, - transitions.observation, - transitions.action, - ).mean(axis=-1) - new_discount = jnp.where( # Inverted done signal - acc_cost - + transitions.extras["state_extras"].get( - "curr_discount", jnp.ones_like(transitions.reward) - ) - * future_cost - > planning_env.safety_budget, - jnp.zeros_like(cost, dtype=jnp.float32), - jnp.ones_like(cost, dtype=jnp.float32), - ) - discount = ( - jnp.logical_or( # Logical OR enforce more ones -> reduce pessimism - transitions.discount, new_discount - ) - ) - else: - discount = transitions.discount next_obs_pred = next_obs_pred.mean(0) return Transition( @@ -267,7 +242,7 @@ def relabel_transitions( next_observation=next_obs_pred, action=transitions.action, reward=new_reward, - discount=discount, + discount=transitions.discount, extras=transitions.extras, ), disagreement diff --git a/ss2r/algorithms/mbpo/types.py b/ss2r/algorithms/mbpo/types.py index 809431ad7..babfbd172 100644 --- a/ss2r/algorithms/mbpo/types.py +++ b/ss2r/algorithms/mbpo/types.py @@ -13,12 +13,12 @@ class TrainingState: policy_params: Params qr_optimizer_state: optax.OptState qr_params: Params - qc_optimizer_state: optax.OptState - qc_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 + 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 53655b85d..22e41dd40 100644 --- a/ss2r/algorithms/sac/q_transforms.py +++ b/ss2r/algorithms/sac/q_transforms.py @@ -42,7 +42,7 @@ def __call__( return target_q -class PessCostUpdate(QTransformation): +class PessimisticCostUpdate(QTransformation): def __call__( self, transitions: Transition, @@ -61,11 +61,8 @@ def __call__( cost * scale + transitions.discount * gamma * next_v ) old_q = q_fn(transitions.observation, transitions.action).mean(axis=-1) - target_q = jax.lax.stop_gradient( - jnp.minimum( - new_target_q, old_q - ) # TODO: check if works (intersection of models) - ) + # TODO: check if works (intersection of models) + target_q = jax.lax.stop_gradient(jnp.minimum(new_target_q, old_q)) return target_q @@ -279,8 +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 == "pess_cost_update": - robustness = PessCostUpdate() + 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/pess_cost_update.yaml b/ss2r/configs/agent/cost_robustness/pess_cost_update.yaml deleted file mode 100644 index f614a125c..000000000 --- a/ss2r/configs/agent/cost_robustness/pess_cost_update.yaml +++ /dev/null @@ -1 +0,0 @@ -name: pess_cost_update \ No newline at end of file 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/experiment/cartpole_mbpo.yaml b/ss2r/configs/experiment/cartpole_mbpo.yaml index ab4a618eb..4989fbbe7 100644 --- a/ss2r/configs/experiment/cartpole_mbpo.yaml +++ b/ss2r/configs/experiment/cartpole_mbpo.yaml @@ -7,6 +7,8 @@ defaults: training: num_timesteps: 150000 num_envs: 10 + safe: false + safety_budget: 100 train_domain_randomization: false eval_domain_randomization: false diff --git a/ss2r/configs/experiment/cartpole_mbpo_safe.yaml b/ss2r/configs/experiment/cartpole_mbpo_safe.yaml deleted file mode 100644 index 5d459d22b..000000000 --- a/ss2r/configs/experiment/cartpole_mbpo_safe.yaml +++ /dev/null @@ -1,21 +0,0 @@ -# @package _global_ -defaults: - - override /environment: cartpole - - override /agent: mbpo - - override /agent/cost_robustness: pess_cost_update - - _self_ - -training: - num_timesteps: 150000 - safe: false - safety_budget: 100 - num_envs: 10 - train_domain_randomization: false - eval_domain_randomization: false - -agent: - min_replay_size: 5000 - sac_batch_size: 512 - critic_grad_updates_per_step: 20 - model_grad_updates_per_step: 25 - pessimism: 0.5 \ No newline at end of file