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
4 changes: 2 additions & 2 deletions ss2r/algorithms/penalizers.py
Original file line number Diff line number Diff line change
Expand Up @@ -251,7 +251,7 @@ def get_penalizer(cfg):
initial_list = []
if cfg.training.get("safe", False):
n_safety_constraints = cfg.agent["model_ensemble_size"]
if cfg.agent.use_mean_critic:
if cfg.agent.use_mean_critic or cfg.agent.use_max_critic:
n_safety_constraints = 1
n_constraints += n_safety_constraints
initial_list.append(
Expand All @@ -278,7 +278,7 @@ def get_penalizer(cfg):
penalty_multiplier_list = []
if cfg.training.get("safe", False):
n_safety_constraints = cfg.agent["model_ensemble_size"]
if cfg.agent.use_mean_critic:
if cfg.agent.use_mean_critic or cfg.agent.use_max_critic:
n_safety_constraints = 1
n_constraints += n_safety_constraints
lagrange_multiplier_list.append(
Expand Down
5 changes: 4 additions & 1 deletion ss2r/algorithms/sbsrl/losses.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@ def make_losses(
safe,
save_sooper_backup,
use_mean_critic,
use_max_critic,
uncertainty_constraint,
uncertainty_epsilon,
n_critics,
Expand Down Expand Up @@ -266,6 +267,8 @@ def actor_loss(
qc_constr = mean_qc
if use_mean_critic:
qc_constr = mean_qc.mean()
if use_max_critic:
qc_constr = mean_qc.max()
aux["qc_max"] = mean_qc.max()
safety_constraint = (
safety_budget - qc_constr
Expand Down Expand Up @@ -293,7 +296,7 @@ def actor_loss(
aux |= penalizer_aux
if safe:
n_safety_constraints = ensemble_size
if use_mean_critic:
if use_mean_critic or use_max_critic:
n_safety_constraints = 1
aux["cost_multipliers"] = penalizer_params.lagrange_multiplier[
:n_safety_constraints
Expand Down
2 changes: 2 additions & 0 deletions ss2r/algorithms/sbsrl/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -238,6 +238,7 @@ def train(
safe: bool = False,
save_sooper_backup: bool = False,
use_mean_critic: bool = False,
use_max_critic: bool = False,
uncertainty_constraint: bool = False,
uncertainty_epsilon: float = 0.0,
safety_budget: float = float("inf"),
Expand Down Expand Up @@ -568,6 +569,7 @@ def normalize_leaf(data: jnp.ndarray, std: jnp.ndarray) -> jnp.ndarray:
safe=safe,
save_sooper_backup=save_sooper_backup,
use_mean_critic=use_mean_critic,
use_max_critic=use_max_critic,
uncertainty_constraint=uncertainty_constraint,
uncertainty_epsilon=uncertainty_epsilon,
n_critics=n_critics,
Expand Down
1 change: 1 addition & 0 deletions ss2r/configs/agent/sbsrl.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -57,5 +57,6 @@ training_step_fn: on_policy
uncertainty_constraint: false
uncertainty_epsilon: 0.0
use_mean_critic: false
use_max_critic: false
load_disagreement_normalizer: false
pessimistic_cost: false
9 changes: 3 additions & 6 deletions ss2r/configs/experiment/cartpole_swingup_simple_sbsrl.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -35,11 +35,8 @@ agent:
uncertainty_epsilon: 0.0
use_mean_critic: false
model_to_real_data_ratio: 0.5
#penalizer:
# initial_multiplier_uncertainty: 0.
# initial_multiplier_cost: 0.
# learning_rate: 0.1
penalizer:
lagrange_multiplier: 0.1
penalty_multiplier: 5e-6
lagrange_multiplier: 5
penalty_multiplier: 1e-6
penalty_multiplier_factor: 1e-4
lagrange_multiplier_sigma: 0
50 changes: 50 additions & 0 deletions ss2r/configs/experiment/humanoid_walk_sbsrl_lagrangian.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
# @package _global_
defaults:
- override /environment: humanoid_walk
- override /agent: sbsrl
- override /agent/data_collection: episodic #TODO: check
- override /agent/cost_robustness: pessimistic_cost_update
- override /agent/penalizer: multiaug_lagrangian
- _self_


training:
num_timesteps: 30000
safe: true
train_domain_randomization: false
eval_domain_randomization: false
safety_budget: 100
num_envs: 1
num_evals: 20
wandb_id: z3518n9m

agent:
policy_hidden_layer_sizes: [256, 256, 256]
value_hidden_layer_sizes: [512, 512]
model_hidden_layer_sizes: [400, 400, 400, 400]
activation: swish
batch_size: 256
min_replay_size: 2000
max_replay_size: 1048576
critic_grad_updates_per_step: 2000
model_grad_updates_per_step: 80000
num_model_rollouts: 100000
learning_rate: 1e-6
critic_learning_rate: 5e-5
num_critic_updates_per_actor_update: 20
model_learning_rate: 1e-4
reward_scaling: 1
cost_scaling: 1
uncertainty_constraint: true
uncertainty_epsilon: 0
model_to_real_data_ratio: 1
use_mean_critic: false
reward_pessimism: 0
cost_pessimism: 0
cost_robustness: null
safety_discounting: 0.999
normalize_budget: false
penalizer:
lagrange_multiplier: 5
penalty_multiplier: 0.001
penalty_multiplier_factor: 1e-4
43 changes: 43 additions & 0 deletions ss2r/configs/experiment/humanoid_walk_sbsrl_offline.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,43 @@
# @package _global_
defaults:
- override /environment: humanoid_walk
- override /agent: sbsrl
- override /agent/penalizer: multiaug_lagrangian
- override /agent/data_collection: step
- _self_

training:
num_timesteps: 2500000
train_domain_randomization: false
eval_domain_randomization: false
safe: true
safety_budget: 100
num_envs: 128
num_evals: 20
wandb_id: 6ldtps0q

agent:
policy_hidden_layer_sizes: [256, 256, 256]
value_hidden_layer_sizes: [512, 512]
model_hidden_layer_sizes: [400, 400, 400, 400]
activation: swish
batch_size: 256
max_replay_size: 4194304
critic_grad_updates_per_step: 10
model_grad_updates_per_step: 50
num_model_rollouts: 10000
discounting: 0.99
safety_discounting: 0.999
normalize_budget: false
reward_scaling: 1
cost_scaling: 1
offline: true
reward_pessimism: 150
uncertainty_constraint: true
uncertainty_epsilon: 0
model_to_real_data_ratio: 0.5
use_mean_critic: false
penalizer:
lagrange_multiplier: 0.1
penalty_multiplier: 0.001
penalty_multiplier_factor: 1e-4
5 changes: 3 additions & 2 deletions ss2r/configs/experiment/humanoid_walk_simple.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -6,12 +6,13 @@ defaults:


training:
num_timesteps: 5000000
num_timesteps: 4000000
safe: true
train_domain_randomization: false
eval_domain_randomization: false
safety_budget: 100
num_eval_episodes: 1

agent:
activation: swish
activation: swish

48 changes: 48 additions & 0 deletions ss2r/configs/experiment/rccar_mbpo_from_sbsrl.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,48 @@
# @package _global_
defaults:
- override /environment: rccar_real
- override /agent: mbpo
- override /agent/penalizer: lagrangian
- _self_

environment:
action_delay: 1
observation_delay: 0
sliding_window: 5
dt: 0.03333333
sample_init_pose: true

training:
num_envs: 1
num_timesteps: 50000
episode_length: 250
safe: true
train_domain_randomization: false
eval_domain_randomization: false
safety_budget: 5.0
wandb_id: null
render: true

agent:
batch_size: 256
min_replay_size: 500
max_replay_size: 1048576
policy_hidden_layer_sizes: [64, 64]
critic_grad_updates_per_step: 10
model_grad_updates_per_step: 50
num_model_rollouts: 100000
learning_rate: 3e-4
critic_learning_rate: 3e-4
model_learning_rate: 3e-4
use_termination: false
pessimism: 30
# Use negative optimism so that the reward is pessimistic w.r.t uncertianty
# following MOPO paper.
optimism: 10
safety_filter: null
offline: false
load_from_sbsrl: true
penalizer:
lagrange_multiplier: 0.1
penalty_multiplier: 0.001
penalty_multiplier_factor: 1e-4
4 changes: 2 additions & 2 deletions ss2r/configs/experiment/rccar_sbsrl_short.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -49,8 +49,8 @@ agent:
#initial_multiplier_uncertainty: 0.01
#initial_multiplier_cost: 10
#learning_rate: 1e-2
lagrange_multiplier: 0.1
penalty_multiplier: 5e-8
lagrange_multiplier: 1
penalty_multiplier: 1e-3
penalty_multiplier_factor: 1e-4


Expand Down