diff --git a/.gitignore b/.gitignore index 080a77b..82c7908 100644 --- a/.gitignore +++ b/.gitignore @@ -12,3 +12,9 @@ datasets/ shells/ ckpts log_bkp/ +scripts/main.py +scripts/rmb_main.py +output.txt +gravity* +*.sh +*.zip \ No newline at end of file diff --git a/scripts/configs/bt_awac/rpl/apl.yaml b/scripts/configs/bt_awac/rpl/apl.yaml new file mode 100644 index 0000000..c3f10d6 --- /dev/null +++ b/scripts/configs/bt_awac/rpl/apl.yaml @@ -0,0 +1,110 @@ +algorithm: + class: BTAWAC + beta: 0.3333 + max_exp_clip: 100.0 + reward_reg: 0.0 + rm_label: true + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +env: walker2d-gravity-1.0 #halfcheetah-gravity-1.0 +env_kwargs: +env_wrapper: +env_wrapper_kwargs: + +eval_env: walker2d-gravity-1.0 +eval_env_kwargs: +eval_env_wrapper: +eval_env_wrapper_kwargs: + +replay: false +label_key: rl_sum + +optim: + default: + class: Adam + lr: 0.0003 + +network: + reward: + class: EnsembleMLP + ensemble_size: 1 + hidden_dims: [256, 256, 256] + reward_act: identity #sigmoid + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256, 256] + reparameterize: false + conditioned_logstd: false + logstd_min: -5 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256, 256] + + +rm_dataset: + - class: RPLComparisonDataset + env: + batch_size: 64 + segment_length: null + label_key: + replay: + odrl: true + +rm_dataloader: + num_workers: 2 + batch_size: null + +rl_dataset: + - class: APLOfflineDataset + env: + batch_size: 256 + replay: + odrl: true + +rl_dataloader: + num_workers: 2 + batch_size: null + +trainer: + env_freq: null + rm_steps: 100000 + rl_steps: 500000 + log_freq: 500 + profile_freq: 500 + eval_freq: 5000 + label_reward: true + normalize_reward: true + +rm_eval: + function: eval_reward_model + eval_dataset_kwargs: + class: RPLComparisonDataset + env: # + batch_size: 32 + label_key: rl_reward_sum + eval: true + replay: + odrl: true + +rl_eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 500000 + +processor: null diff --git a/scripts/configs/bt_awac/rpl/multi.yaml b/scripts/configs/bt_awac/rpl/multi.yaml new file mode 100644 index 0000000..5f7371d --- /dev/null +++ b/scripts/configs/bt_awac/rpl/multi.yaml @@ -0,0 +1,113 @@ +algorithm: + class: BTAWAC + beta: 0.3333 + max_exp_clip: 100.0 + reward_reg: 0.0 + rm_label: true + +num_tasks: 5 +task_name: 0.1_0.5_1.0_2.0_5.0 + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +env: halfcheetah-gravity-1.0 +env_kwargs: +env_wrapper: +env_wrapper_kwargs: + +# eval_env: halfcheetah-gravity-1.0 +# eval_env_kwargs: +# eval_env_wrapper: +# eval_env_wrapper_kwargs: + +replay: true +label_key: rl_dir + +optim: + default: + class: Adam + lr: 0.0003 + +network: + reward: + class: EnsembleMLP + ensemble_size: 1 + hidden_dims: [256, 256, 256] + reward_act: sigmoid + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256, 256] + reparameterize: false + conditioned_logstd: false + logstd_min: -5 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256, 256] + + +rm_dataset: + - class: MultiRPLComparisonDataset + label_key: + num_tasks: + env: + task_name: + batch_size: 64 + odrl: true + +rm_dataloader: + num_workers: 2 + batch_size: null + +rl_dataset: + - class: RPLOfflineDataset + env: + batch_size: 256 + replay: + odrl: true + +rl_dataloader: + num_workers: 2 + batch_size: null + +trainer: + env_freq: null + rm_steps: 100000 + rl_steps: 500000 + log_freq: 500 + profile_freq: 500 + eval_freq: 10000 + label_reward: true + normalize_reward: true + +rm_eval: + function: eval_reward_model + eval_dataset_kwargs: + class: RPLComparisonDataset + env: # + batch_size: 32 + label_key: rl_reward_sum + eval: true + replay: + odrl: true + +rl_eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 500000 + +processor: null diff --git a/scripts/configs/bt_awac/rpl/multi_d4rl.yaml b/scripts/configs/bt_awac/rpl/multi_d4rl.yaml new file mode 100644 index 0000000..1679178 --- /dev/null +++ b/scripts/configs/bt_awac/rpl/multi_d4rl.yaml @@ -0,0 +1,111 @@ +algorithm: + class: BTAWAC + beta: 0.3333 + max_exp_clip: 100.0 + reward_reg: 0.0 + rm_label: true + +num_tasks: 5 +task_name: 0.1_0.5_1.0_2.0_5.0 + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +env: halfcheetah-gravity-1.0 +env_kwargs: +env_wrapper: +env_wrapper_kwargs: + +# eval_env: halfcheetah-gravity-1.0 +# eval_env_kwargs: +# eval_env_wrapper: +# eval_env_wrapper_kwargs: + +replay: true +label_key: rl_dir + +optim: + default: + class: Adam + lr: 0.0003 + +network: + reward: + class: EnsembleMLP + ensemble_size: 1 + hidden_dims: [256, 256, 256] + reward_act: sigmoid + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256, 256] + reparameterize: false + conditioned_logstd: false + logstd_min: -5 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256, 256] + + +rm_dataset: + - class: MultiRPLComparisonDataset + label_key: + num_tasks: + env: + task_name: + batch_size: 64 + odrl: true + +rm_dataloader: + num_workers: 2 + batch_size: null + +rl_dataset: + - class: D4RLOfflineDataset + env: + batch_size: 256 + +rl_dataloader: + num_workers: 2 + batch_size: null + +trainer: + env_freq: null + rm_steps: 100000 + rl_steps: 500000 + log_freq: 500 + profile_freq: 500 + eval_freq: 10000 + label_reward: true + normalize_reward: true + +rm_eval: + function: eval_reward_model + eval_dataset_kwargs: + class: RPLComparisonDataset + env: # + batch_size: 32 + label_key: rl_reward_sum + eval: true + replay: + odrl: true + +rl_eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 500000 + +processor: null diff --git a/scripts/configs/bt_awac/rpl/odrl.yaml b/scripts/configs/bt_awac/rpl/odrl.yaml new file mode 100644 index 0000000..55bf833 --- /dev/null +++ b/scripts/configs/bt_awac/rpl/odrl.yaml @@ -0,0 +1,110 @@ +algorithm: + class: BTAWAC + beta: 0.3333 + max_exp_clip: 100.0 + reward_reg: 0.0 + rm_label: true + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +env: halfcheetah-gravity-0.5 +env_kwargs: +env_wrapper: +env_wrapper_kwargs: + +eval_env: halfcheetah-gravity-1.0 +eval_env_kwargs: +eval_env_wrapper: +eval_env_wrapper_kwargs: + +replay: true +label_key: rl_dir + +optim: + default: + class: Adam + lr: 0.0003 + +network: + reward: + class: EnsembleMLP + ensemble_size: 1 + hidden_dims: [256, 256, 256] + reward_act: identity #sigmoid + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256, 256] + reparameterize: false + conditioned_logstd: false + logstd_min: -5 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256, 256] + + +rm_dataset: + - class: RPLComparisonDataset + env: + batch_size: 64 + segment_length: null + label_key: + replay: + odrl: true + +rm_dataloader: + num_workers: 2 + batch_size: null + +rl_dataset: + - class: RPLOfflineDataset + env: + batch_size: 256 + replay: + odrl: true + +rl_dataloader: + num_workers: 2 + batch_size: null + +trainer: + env_freq: null + rm_steps: 100000 + rl_steps: 500000 + log_freq: 500 + profile_freq: 500 + eval_freq: 10000 + label_reward: true + normalize_reward: true + +rm_eval: + function: eval_reward_model + eval_dataset_kwargs: + class: RPLComparisonDataset + env: # + batch_size: 32 + label_key: rl_reward_sum + eval: true + replay: + odrl: true + +rl_eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 500000 + +processor: null diff --git a/scripts/configs/bt_awac/rpl/odrl_d4rl.yaml b/scripts/configs/bt_awac/rpl/odrl_d4rl.yaml new file mode 100644 index 0000000..ea3d32e --- /dev/null +++ b/scripts/configs/bt_awac/rpl/odrl_d4rl.yaml @@ -0,0 +1,108 @@ +algorithm: + class: BTAWAC + beta: 0.3333 + max_exp_clip: 100.0 + reward_reg: 0.0 + rm_label: true + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +env: halfcheetah-gravity-0.5 +env_kwargs: +env_wrapper: +env_wrapper_kwargs: + +eval_env: halfcheetah-gravity-1.0 +eval_env_kwargs: +eval_env_wrapper: +eval_env_wrapper_kwargs: + +replay: true +label_key: rl_dir + +optim: + default: + class: Adam + lr: 0.0003 + +network: + reward: + class: EnsembleMLP + ensemble_size: 1 + hidden_dims: [256, 256, 256] + reward_act: sigmoid + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256, 256] + reparameterize: false + conditioned_logstd: false + logstd_min: -5 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256, 256] + + +rm_dataset: + - class: RPLComparisonDataset + env: + batch_size: 64 + segment_length: null + label_key: + replay: + odrl: true + +rm_dataloader: + num_workers: 2 + batch_size: null + +rl_dataset: + - class: D4RLOfflineDataset + env: + batch_size: 256 + +rl_dataloader: + num_workers: 2 + batch_size: null + +trainer: + env_freq: null + rm_steps: 100000 + rl_steps: 500000 + log_freq: 500 + profile_freq: 500 + eval_freq: 10000 + label_reward: true + normalize_reward: true + +rm_eval: + function: eval_reward_model + eval_dataset_kwargs: + class: RPLComparisonDataset + env: # + batch_size: 32 + label_key: rl_reward_sum + eval: true + replay: + odrl: true + +rl_eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 500000 + +processor: null diff --git a/scripts/configs/bt_awac/rpl/rpl.yaml b/scripts/configs/bt_awac/rpl/rpl.yaml new file mode 100644 index 0000000..7d1cf4e --- /dev/null +++ b/scripts/configs/bt_awac/rpl/rpl.yaml @@ -0,0 +1,119 @@ +algorithm: + class: BTAWAC + beta: 0.3333 + max_exp_clip: 100.0 + reward_reg: 0.0 + rm_label: true + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +src_gravity: 0.5 +src_variant: gravity-50 +tgt_gravity: 1.5 +tgt_variant: gravity-150 +replay: false + +env: HalfCheetah-v3 +env_kwargs: +env_wrapper: MujocoParamOverWrite +env_wrapper_kwargs: + overwrite_args: + gravity: # gravity-150 + do_scale: True + +eval_env: +eval_env_kwargs: +eval_env_wrapper: MujocoParamOverWrite +eval_env_wrapper_kwargs: + overwrite_args: + gravity: # gravity-150 + do_scale: True + +optim: + default: + class: Adam + lr: 0.0003 + +network: + reward: + class: EnsembleMLP + ensemble_size: 1 + hidden_dims: [256, 256] + reward_act: identity #sigmoid + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256] + reparameterize: false + conditioned_logstd: false + logstd_min: -5 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256] + + +rm_dataset: + - class: RPLComparisonDataset + env: + batch_size: 8 + segment_length: null + label_key: rl_sum + variant: + replay: + +rm_dataloader: + num_workers: 2 + batch_size: null + +rl_dataset: + - class: RPLOfflineDataset + env: + batch_size: 256 + variant: + replay: + +rl_dataloader: + num_workers: 2 + batch_size: null + +trainer: + env_freq: null + rm_label: true + rm_steps: 1000000 + rl_steps: 500000 + log_freq: 500 + profile_freq: 500 + eval_freq: 5000 + save_rm_path: + +rm_eval: + function: eval_reward_model + eval_dataset_kwargs: + class: RPLComparisonDataset + env: + batch_size: 32 + label_key: rl_sum + variant: + eval: true + replay: + +rl_eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 1000000 + +processor: null diff --git a/scripts/configs/bt_iql/gym/default.yaml b/scripts/configs/bt_iql/gym/default.yaml index 8184f90..e3894c5 100644 --- a/scripts/configs/bt_iql/gym/default.yaml +++ b/scripts/configs/bt_iql/gym/default.yaml @@ -49,7 +49,7 @@ network: hidden_dims: [256, 256] rm_dataset: - - class: IPLComparisonOfflineDataset + - class: RPLComparisonOfflineDataset env: batch_size: 8 segment_length: null diff --git a/scripts/configs/bt_iql/rpl/rpl.yaml b/scripts/configs/bt_iql/rpl/rpl.yaml new file mode 100644 index 0000000..41baa17 --- /dev/null +++ b/scripts/configs/bt_iql/rpl/rpl.yaml @@ -0,0 +1,122 @@ +algorithm: + class: BTIQL + beta: 0.3333 + expectile: 0.7 + max_exp_clip: 100.0 + reward_reg: 0.0 + rm_label: true + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +src_gravity: 0.5 +src_variant: gravity-50 +tgt_gravity: 1.5 +tgt_variant: gravity-150 +replay: false + +env: HalfCheetah-v3 +env_kwargs: +env_wrapper: MujocoParamOverWrite +env_wrapper_kwargs: + overwrite_args: + gravity: # gravity-150 + do_scale: True + +eval_env: +eval_env_kwargs: +eval_env_wrapper: MujocoParamOverWrite +eval_env_wrapper_kwargs: + overwrite_args: + gravity: # gravity-150 + do_scale: True + +optim: + default: + class: Adam + lr: 0.0003 + +network: + reward: + class: EnsembleMLP + ensemble_size: 1 + hidden_dims: [256, 256] + reward_act: sigmoid + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256] + reparameterize: false + conditioned_logstd: false + logstd_min: -5 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256] + value: + class: Critic + ensemble_size: 1 + hidden_dims: [256, 256] + +rm_dataset: + - class: RPLComparisonDataset + env: + batch_size: 8 + segment_length: null + label_key: rl_sum + variant: + replay: + +rm_dataloader: + num_workers: 2 + batch_size: null + +rl_dataset: + - class: RPLOfflineDataset + env: + batch_size: 256 + variant: + replay: + +rl_dataloader: + num_workers: 2 + batch_size: null + +trainer: + env_freq: null + rm_label: true + rm_steps: 100000 + rl_steps: 1000000 + log_freq: 500 + profile_freq: 500 + eval_freq: 5000 + +rm_eval: + function: eval_reward_model + eval_dataset_kwargs: + class: RPLComparisonDataset + env: + batch_size: 32 + label_key: rl_sum + variant: + eval: true + replay: + +rl_eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 1000000 + +processor: null diff --git a/scripts/configs/bt_iql/rpl/rpl_multi.yaml b/scripts/configs/bt_iql/rpl/rpl_multi.yaml new file mode 100644 index 0000000..ab8f6b0 --- /dev/null +++ b/scripts/configs/bt_iql/rpl/rpl_multi.yaml @@ -0,0 +1,118 @@ +algorithm: + class: BTIQL + beta: 0.3333 + expectile: 0.7 + max_exp_clip: 100.0 + reward_reg: 0.0 + rm_label: true + +num_tasks: 5 +task_name: 0.1_0.5_1.0_2.0_5.0 + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +env: halfcheetah-gravity-1.0 +env_kwargs: +env_wrapper: +env_wrapper_kwargs: + +# eval_env: halfcheetah-gravity-0.1 +# eval_env_kwargs: +# eval_env_wrapper: +# eval_env_wrapper_kwargs: + +replay: true +label_key: rl_sum + +optim: + default: + class: Adam + lr: 0.0003 + +network: + reward: + class: EnsembleMLP + ensemble_size: 1 + hidden_dims: [256, 256] + reward_act: identity #sigmoid + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256] + reparameterize: false + conditioned_logstd: false + logstd_min: -5 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256] + value: + class: Critic + ensemble_size: 1 + hidden_dims: [256, 256] + + +rm_dataset: + - class: MultiRPLComparisonDataset # please fill in the class and parameters here + label_key: + num_tasks: + env: + task_name: + batch_size: 32 #8 + variant: gravity-100 + odrl: true + +rm_dataloader: + num_workers: 2 + batch_size: null + +rl_dataset: + - class: RPLOfflineDataset + env: + batch_size: 256 + replay: + odrl: true + +rl_dataloader: + num_workers: 2 + batch_size: null + +trainer: + env_freq: null + rm_label: true + rm_steps: 100000 + rl_steps: 1000000 + log_freq: 500 + profile_freq: 500 + eval_freq: 5000 + +rm_eval: + function: eval_reward_model + eval_dataset_kwargs: + class: RPLComparisonDataset + env: # + batch_size: 32 + label_key: rl_reward_sum + eval: true + replay: + odrl: true + +rl_eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 1000000 + +processor: null diff --git a/scripts/configs/bt_iql/rpl/rpl_odrl.yaml b/scripts/configs/bt_iql/rpl/rpl_odrl.yaml new file mode 100644 index 0000000..a0a4e8c --- /dev/null +++ b/scripts/configs/bt_iql/rpl/rpl_odrl.yaml @@ -0,0 +1,114 @@ +algorithm: + class: BTIQL + beta: 0.3333 + expectile: 0.7 + max_exp_clip: 100.0 + reward_reg: 0.0 + rm_label: true + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +env: halfcheetah-gravity-5.0 +env_kwargs: +env_wrapper: +env_wrapper_kwargs: + +eval_env: halfcheetah-gravity-0.1 +eval_env_kwargs: +eval_env_wrapper: +eval_env_wrapper_kwargs: + +replay: true +label_key: rl_sum + +optim: + default: + class: Adam + lr: 0.0003 + +network: + reward: + class: EnsembleMLP + ensemble_size: 1 + hidden_dims: [256, 256] + reward_act: identity #sigmoid + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256] + reparameterize: false + conditioned_logstd: false + logstd_min: -5 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256] + value: + class: Critic + ensemble_size: 1 + hidden_dims: [256, 256] + + +rm_dataset: + - class: RPLComparisonDataset + env: + batch_size: 8 + segment_length: null + label_key: + replay: + odrl: true + +rm_dataloader: + num_workers: 2 + batch_size: null + +rl_dataset: + - class: RPLOfflineDataset + env: + batch_size: 256 + replay: + odrl: true + +rl_dataloader: + num_workers: 2 + batch_size: null + +trainer: + env_freq: null + rm_label: true + rm_steps: 100000 + rl_steps: 1000000 + log_freq: 500 + profile_freq: 500 + eval_freq: 5000 + +rm_eval: + function: eval_reward_model + eval_dataset_kwargs: + class: RPLComparisonDataset + env: # + batch_size: 32 + label_key: + eval: true + replay: + odrl: true + +rl_eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 1000000 + +processor: null diff --git a/scripts/configs/bt_sac/rpl/rpl.yaml b/scripts/configs/bt_sac/rpl/rpl.yaml new file mode 100644 index 0000000..9e62ab0 --- /dev/null +++ b/scripts/configs/bt_sac/rpl/rpl.yaml @@ -0,0 +1,115 @@ +algorithm: + class: BTSAC + alpha: 0.2 + auto_alpha: false #true + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +src_gravity: 1.0 +src_variant: gravity-100 +tgt_gravity: 1.0 +tgt_variant: gravity-100 +replay: false + +env: HalfCheetah-v3 +env_kwargs: +env_wrapper: MujocoParamOverWrite +env_wrapper_kwargs: + overwrite_args: + gravity: # gravity-150 + do_scale: True + +eval_env: +eval_env_kwargs: +eval_env_wrapper: MujocoParamOverWrite +eval_env_wrapper_kwargs: + overwrite_args: + gravity: # gravity-150 + do_scale: True + +optim: + default: + class: Adam + lr: 0.0003 + +network: + reward: + class: EnsembleMLP + ensemble_size: 1 + hidden_dims: [256, 256] + reward_act: identity #sigmoid + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256] + logstd_min: -20 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256] + + +rm_dataset: + - class: RPLComparisonDataset + env: + batch_size: 8 + segment_length: null + label_key: rl_sum + variant: + replay: + +rm_dataloader: + num_workers: 2 + batch_size: null + +rl_dataset: + - class: RPLOfflineDataset + env: + batch_size: 256 + variant: + replay: + +rl_dataloader: + num_workers: 2 + batch_size: null + +trainer: + env_freq: null + rm_label: true + rm_steps: 1000000 + rl_steps: 500000 + log_freq: 500 + profile_freq: 500 + eval_freq: 5000 + #load_rm_path: + +rm_eval: + function: eval_reward_model + eval_dataset_kwargs: + class: RPLComparisonDataset + env: + batch_size: 32 + label_key: rl_sum + variant: + eval: true + replay: + +rl_eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 1000000 + +processor: null diff --git a/scripts/configs/bt_sac_online/rpl/rpl.yaml b/scripts/configs/bt_sac_online/rpl/rpl.yaml new file mode 100644 index 0000000..de2e8a6 --- /dev/null +++ b/scripts/configs/bt_sac_online/rpl/rpl.yaml @@ -0,0 +1,111 @@ +algorithm: + class: BTSAC + alpha: 0.2 + auto_alpha: false #true + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +src_gravity: 1.0 +src_variant: gravity-100 +tgt_gravity: 1.0 +tgt_variant: gravity-100 +replay: false + +env: HalfCheetah-v3 +env_kwargs: +env_wrapper: MujocoParamOverWrite +env_wrapper_kwargs: + overwrite_args: + gravity: # gravity-150 + do_scale: True + +eval_env: +eval_env_kwargs: +eval_env_wrapper: MujocoParamOverWrite +eval_env_wrapper_kwargs: + overwrite_args: + gravity: # gravity-150 + do_scale: True + +optim: + default: + class: Adam + lr: 0.0003 + +network: + reward: + class: EnsembleMLP + ensemble_size: 1 + hidden_dims: [256, 256] + reward_act: identity #sigmoid + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256] + logstd_min: -20 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256] + + +rm_dataset: + - class: RPLComparisonDataset + env: + batch_size: 8 + segment_length: null + label_key: rl_sum + variant: + replay: + +rm_dataloader: + num_workers: 2 + batch_size: null + +buffer: + max_buffer_size: 1000000 + batch_size: 256 + +trainer: + random_policy_step: 1000 #5000 + warmup_step: 2000 + max_trajectory_length: 1000 + env_freq: null + rm_label: true + rm_steps: 1000000 + rl_steps: 500000 + log_freq: 500 + profile_freq: 500 + eval_freq: 5000 + load_rm_path: + +rm_eval: + function: eval_reward_model + eval_dataset_kwargs: + class: RPLComparisonDataset + env: + batch_size: 32 + label_key: rl_sum + variant: + eval: true + replay: + +rl_eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 1000000 + +processor: null diff --git a/scripts/configs/bt_td3bc/rpl/rpl.yaml b/scripts/configs/bt_td3bc/rpl/rpl.yaml new file mode 100644 index 0000000..e01e246 --- /dev/null +++ b/scripts/configs/bt_td3bc/rpl/rpl.yaml @@ -0,0 +1,118 @@ +algorithm: + class: BTTD3BC + alpha: 0.2 + policy_noise: 0.2 + noise_clip: 0.5 + max_action: 1.0 + discount: 0.99 + tau: 0.005 + actor_update_interval: 2 + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +src_gravity: 1.0 +src_variant: gravity-100 +tgt_gravity: 1.0 +tgt_variant: gravity-100 +replay: false + +env: HalfCheetah-v3 +env_kwargs: +env_wrapper: MujocoParamOverWrite +env_wrapper_kwargs: + overwrite_args: + gravity: # gravity-150 + do_scale: True + +eval_env: +eval_env_kwargs: +eval_env_wrapper: MujocoParamOverWrite +eval_env_wrapper_kwargs: + overwrite_args: + gravity: # gravity-150 + do_scale: True + +optim: + default: + class: Adam + lr: 0.0003 + +network: + reward: + class: EnsembleMLP + ensemble_size: 1 + hidden_dims: [256, 256] + reward_act: identity #sigmoid + actor: + class: SquashedDeterministicActor + hidden_dims: [256, 256] + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256] + + +rm_dataset: + - class: RPLComparisonDataset + env: + batch_size: 8 + segment_length: null + label_key: rl_sum + variant: + replay: + +rm_dataloader: + num_workers: 2 + batch_size: null + +rl_dataset: + - class: RPLOfflineDataset + env: + batch_size: 256 + variant: + replay: + +rl_dataloader: + num_workers: 2 + batch_size: null + +trainer: + env_freq: null + rm_label: true + rm_steps: 1000000 + rl_steps: 500000 + log_freq: 500 + profile_freq: 500 + eval_freq: 5000 + load_rm_path: + +rm_eval: + function: eval_reward_model + eval_dataset_kwargs: + class: RPLComparisonDataset + env: + batch_size: 32 + label_key: rl_sum + variant: + eval: true + replay: + +rl_eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 1000000 + +processor: null diff --git a/scripts/configs/oracle_awac/rpl/apl.yaml b/scripts/configs/oracle_awac/rpl/apl.yaml new file mode 100644 index 0000000..61e6b1e --- /dev/null +++ b/scripts/configs/oracle_awac/rpl/apl.yaml @@ -0,0 +1,79 @@ +algorithm: + class: OracleAWAC + beta: 0.3333 + max_exp_clip: 100.0 + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +env: hopper-gravity-1.0 #walker2d-gravity-1.0 #halfcheetah-gravity-1.0 +env_kwargs: +env_wrapper: +env_wrapper_kwargs: + +eval_env: hopper-gravity-1.0 #walker2d-gravity-1.0 #halfcheetah-gravity-1.0 +eval_env_kwargs: +eval_env_wrapper: +eval_env_wrapper_kwargs: + +replay: false + +mismatch: hopper-gravity-0.5 #walker2d-gravity-5.0 #halfcheetah-gravity-1.0 + +optim: + default: + class: Adam + lr: 0.0003 + +network: + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256, 256] + reparameterize: false + conditioned_logstd: false + logstd_min: -5 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256, 256] + +dataset: + - class: APLOfflineDataset + env: + batch_size: 256 + replay: + odrl: true + mismatch: + +dataloader: + num_workers: 2 + batch_size: null + + +trainer: + env_freq: null + total_steps: 500000 + log_freq: 500 + profile_freq: 500 + eval_freq: 5000 + normalize_reward: true + +eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 500000 + +processor: null diff --git a/scripts/configs/oracle_awac/rpl/odrl.yaml b/scripts/configs/oracle_awac/rpl/odrl.yaml new file mode 100644 index 0000000..63d994c --- /dev/null +++ b/scripts/configs/oracle_awac/rpl/odrl.yaml @@ -0,0 +1,76 @@ +algorithm: + class: OracleAWAC + beta: 0.3333 + max_exp_clip: 100.0 + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +env: halfcheetah-gravity-1.0 +env_kwargs: +env_wrapper: +env_wrapper_kwargs: + +eval_env: halfcheetah-gravity-1.0 +eval_env_kwargs: +eval_env_wrapper: +eval_env_wrapper_kwargs: + +replay: true + +optim: + default: + class: Adam + lr: 0.0003 + +network: + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256, 256] + reparameterize: false + conditioned_logstd: false + logstd_min: -5 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256, 256] + +dataset: + - class: RPLOfflineDataset + env: + batch_size: 256 + replay: + odrl: true + +dataloader: + num_workers: 2 + batch_size: null + + +trainer: + env_freq: null + total_steps: 500000 + log_freq: 500 + profile_freq: 500 + eval_freq: 5000 + normalize_reward: true + +eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 500000 + +processor: null diff --git a/scripts/configs/oracle_awac/rpl/rpl.yaml b/scripts/configs/oracle_awac/rpl/rpl.yaml new file mode 100644 index 0000000..9fb8c03 --- /dev/null +++ b/scripts/configs/oracle_awac/rpl/rpl.yaml @@ -0,0 +1,73 @@ +algorithm: + class: OracleAWAC + beta: 0.3333 + max_exp_clip: 100.0 + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +gravity: 1.0 +variant: gravity-100 + +env: HalfCheetah-v3 +env_kwargs: +env_wrapper: MujocoParamOverWrite +env_wrapper_kwargs: + overwrite_args: + gravity: + do_scale: True + +optim: + default: + class: Adam + lr: 0.0003 + +network: + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256] + reparameterize: false + conditioned_logstd: false + logstd_min: -5 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256] + +dataset: + - class: RPLOfflineDataset + env: + batch_size: 256 + variant: + +dataloader: + num_workers: 2 + batch_size: null + + +trainer: + env_freq: null + total_steps: 500000 + log_freq: 500 + profile_freq: 500 + eval_freq: 10000 + +eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 500000 + +processor: null diff --git a/scripts/configs/oracle_iql/rpl/rpl.yaml b/scripts/configs/oracle_iql/rpl/rpl.yaml new file mode 100644 index 0000000..bce9b25 --- /dev/null +++ b/scripts/configs/oracle_iql/rpl/rpl.yaml @@ -0,0 +1,78 @@ +algorithm: + class: OracleIQL + beta: 0.3333 + expectile: 0.7 + max_exp_clip: 100.0 + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +gravity: 1.0 +variant: gravity-100 + +env: HalfCheetah-v3 +env_kwargs: +env_wrapper: MujocoParamOverWrite +env_wrapper_kwargs: + overwrite_args: + gravity: + do_scale: True + +optim: + default: + class: Adam + lr: 0.0003 + +network: + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256] + reparameterize: false + conditioned_logstd: false + logstd_min: -5 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256] + value: + class: Critic + ensemble_size: 1 + hidden_dims: [256, 256] + +dataset: + - class: RPLOfflineDataset + env: + batch_size: 256 + variant: gravity-100 + capacity: 5000 +dataloader: + num_workers: 2 + batch_size: null + + +trainer: + env_freq: null + total_steps: 1000000 + log_freq: 500 + profile_freq: 500 + eval_freq: 5000 + +eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 500000 + +processor: null diff --git a/scripts/configs/oracle_sac/rpl/rpl.yaml b/scripts/configs/oracle_sac/rpl/rpl.yaml new file mode 100644 index 0000000..8183d5b --- /dev/null +++ b/scripts/configs/oracle_sac/rpl/rpl.yaml @@ -0,0 +1,71 @@ +algorithm: + class: OracleSAC + alpha: 0.2 + auto_alpha: false #true + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +gravity: 1.0 +variant: gravity-100 + +env: HalfCheetah-v3 +env_kwargs: +env_wrapper: MujocoParamOverWrite +env_wrapper_kwargs: + overwrite_args: + gravity: + do_scale: True + +optim: + default: + class: Adam + lr: 0.0003 + +network: + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256] + logstd_min: -20 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256] + +dataset: + - class: RPLOfflineDataset + env: + batch_size: 256 + variant: + +dataloader: + num_workers: 2 + batch_size: null + + +trainer: + env_freq: null + total_steps: 1000000 + log_freq: 500 + profile_freq: 500 + eval_freq: 5000 + +eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 500000 + +processor: null diff --git a/scripts/configs/oracle_sac/rpl/rpl1.yaml b/scripts/configs/oracle_sac/rpl/rpl1.yaml new file mode 100644 index 0000000..8ecc71b --- /dev/null +++ b/scripts/configs/oracle_sac/rpl/rpl1.yaml @@ -0,0 +1,60 @@ +algorithm: + alpha: 0.2 + class: OracleSAC +checkpoint: null +dataloader: + batch_size: null + num_workers: 2 +dataset: +- batch_size: 256 + class: RPLOfflineDataset + env: HalfCheetah-v3 + variant: gravity-100 +debug: false +device: null +env: HalfCheetah-v3 +env_kwargs: null +env_wrapper: MujocoParamOverWrite +env_wrapper_kwargs: + do_scale: true + overwrite_args: + gravity: 1.0 +eval: + deterministic: true + function: eval_offline + num_ep: 10 +gravity: 1.0 +name: default +network: + actor: + class: SquashedGaussianActor + hidden_dims: + - 256 + - 256 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: + - 256 + - 256 +optim: + default: + class: Adam + lr: 0.0003 +processor: null +schedulers: + actor: + T_max: 500000 + class: CosineAnnealingLR +seed: 0 +trainer: + env_freq: null + eval_freq: 5000 + log_freq: 500 + profile_freq: 500 + total_steps: 1000000 +variant: gravity-100 +wandb: + activate: false + entity: null + project: null diff --git a/scripts/configs/oracle_sac_online/rpl/rpl.yaml b/scripts/configs/oracle_sac_online/rpl/rpl.yaml new file mode 100644 index 0000000..0b47f26 --- /dev/null +++ b/scripts/configs/oracle_sac_online/rpl/rpl.yaml @@ -0,0 +1,67 @@ +algorithm: + class: OracleSAC + alpha: 0.2 + auto_alpha: true + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +gravity: 1.0 +variant: gravity-100 + +env: HalfCheetah-v3 +env_kwargs: +env_wrapper: MujocoParamOverWrite +env_wrapper_kwargs: + overwrite_args: + gravity: + do_scale: True + +optim: + default: + class: Adam + lr: 0.0003 + +network: + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256] + logstd_min: -20 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256] + +buffer: + max_buffer_size: 1000000 + batch_size: 256 + +trainer: + random_policy_step: 5000 #5000 + warmup_step: 2000 + max_trajectory_length: 1000 + env_freq: null + total_steps: 1000000 + log_freq: 500 + profile_freq: 500 + eval_freq: 5000 + +eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 500000 + +processor: null diff --git a/scripts/configs/oracle_sac_online/rpl/rpl1.yaml b/scripts/configs/oracle_sac_online/rpl/rpl1.yaml new file mode 100644 index 0000000..a5c3edb --- /dev/null +++ b/scripts/configs/oracle_sac_online/rpl/rpl1.yaml @@ -0,0 +1,65 @@ +algorithm: + class: OracleSAC + alpha: 0.2 + auto_alpha: false + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +gravity: 1.0 +variant: gravity-100 + +env: HalfCheetah-v3 +env_kwargs: +env_wrapper: MujocoParamOverWrite +env_wrapper_kwargs: + overwrite_args: + gravity: + do_scale: True + +optim: + default: + class: Adam + lr: 0.0003 + +network: + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256] + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256] + +buffer: + max_buffer_size: 100000 + batch_size: 256 + +trainer: + random_policy_step: 1000 #5000 + warmup_step: 2000 + max_trajectory_length: 1000 + env_freq: null + total_steps: 1000000 + log_freq: 500 + profile_freq: 500 + eval_freq: 5000 + +eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 500000 + +processor: null diff --git a/scripts/configs/oracle_td3bc/rpl/rpl.yaml b/scripts/configs/oracle_td3bc/rpl/rpl.yaml new file mode 100644 index 0000000..7aff44b --- /dev/null +++ b/scripts/configs/oracle_td3bc/rpl/rpl.yaml @@ -0,0 +1,74 @@ +algorithm: + class: OracleTD3BC + alpha: 0.2 + policy_noise: 0.2 + noise_clip: 0.5 + max_action: 1.0 + discount: 0.99 + tau: 0.005 + actor_update_interval: 2 + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +gravity: 1.0 +variant: gravity-100 + +env: HalfCheetah-v3 +env_kwargs: +env_wrapper: MujocoParamOverWrite +env_wrapper_kwargs: + overwrite_args: + gravity: + do_scale: True + +optim: + default: + class: Adam + lr: 0.0003 + +network: + actor: + class: SquashedDeterministicActor + hidden_dims: [256, 256] + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256] + +dataset: + - class: RPLOfflineDataset + env: + batch_size: 256 + variant: + +dataloader: + num_workers: 2 + batch_size: null + + +trainer: + env_freq: null + total_steps: 1000000 + log_freq: 500 + profile_freq: 500 + eval_freq: 5000 + +eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 500000 + +processor: null diff --git a/scripts/configs/rpl/rpl_awac/d4rl.yaml b/scripts/configs/rpl/rpl_awac/d4rl.yaml new file mode 100644 index 0000000..11d44d0 --- /dev/null +++ b/scripts/configs/rpl/rpl_awac/d4rl.yaml @@ -0,0 +1,117 @@ +num_tasks: 5 +task_name: 0.1_0.5_1.0_2.0_5.0 + +algorithm: + class: RPL_AWAC + num_tasks: # number of tasks + alpha: 0.7 # expectile for Eq. 8 + beta: 0.3333 # inv. temperature of IQL + max_exp_clip: 100.0 + rm_label: true + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +env: walker2d-gravity-1.0 +env_kwargs: +env_wrapper: +env_wrapper_kwargs: + +replay: true +label_key: rl_dir + +optim: + default: + class: Adam + lr: 0.0003 + +network: + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256, 256] + reparameterize: false + conditioned_logstd: false + logstd_min: -5 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256, 256] + value: + class: Critic + ensemble_size: 1 + hidden_dims: [256, 256, 256] + reward: + class: EnsembleMLP + ensemble_size: 1 + hidden_dims: [256, 256, 256] + reward_act: sigmoid + optimal: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256, 256] + +rm_dataset: + - class: MultiRPLComparisonDataset # please fill in the class and parameters here + label_key: + num_tasks: + env: + task_name: + batch_size: 64 + variant: gravity-100 + odrl: true +rm_dataloader: + num_workers: 2 + batch_size: null + +rl_dataset: + - class: D4RLOfflineDataset + env: + batch_size: 256 + #replay: + #odrl: true + +rl_dataloader: + num_workers: 2 + batch_size: null + + +trainer: + env_freq: null + rm_steps: 500000 + rl_steps: 500000 + log_freq: 500 + profile_freq: 500 + eval_freq: 10000 + label_reward: true + normalize_reward: true + +rm_eval: + function: eval_reward_model + eval_dataset_kwargs: + class: RPLComparisonDataset + env: + batch_size: 32 + label_key: rl_reward_sum + eval: true + replay: + odrl: true + +rl_eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 500000 + +processor: null diff --git a/scripts/configs/rpl/rpl_awac/odrl.yaml b/scripts/configs/rpl/rpl_awac/odrl.yaml new file mode 100644 index 0000000..f8c4e7f --- /dev/null +++ b/scripts/configs/rpl/rpl_awac/odrl.yaml @@ -0,0 +1,117 @@ +num_tasks: 4 +task_name: 0.1_0.5_2.0_5.0 + +algorithm: + class: RPL_AWAC + num_tasks: # number of tasks + alpha: 0.7 # expectile for Eq. 8 + beta: 0.3333 # inv. temperature of IQL + max_exp_clip: 100.0 + rm_label: true + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +env: walker2d-gravity-1.0 +env_kwargs: +env_wrapper: +env_wrapper_kwargs: + +replay: true +label_key: rl_dir + +optim: + default: + class: Adam + lr: 0.0003 + +network: + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256, 256] + reparameterize: false + conditioned_logstd: false + logstd_min: -5 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256, 256] + value: + class: Critic + ensemble_size: 1 + hidden_dims: [256, 256, 256] + reward: + class: EnsembleMLP + ensemble_size: 1 + hidden_dims: [256, 256, 256] + reward_act: sigmoid + optimal: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256, 256] + +rm_dataset: + - class: MultiRPLComparisonDataset # please fill in the class and parameters here + label_key: + num_tasks: + env: + task_name: + batch_size: 64 + variant: gravity-100 + odrl: true +rm_dataloader: + num_workers: 2 + batch_size: null + +rl_dataset: + - class: RPLOfflineDataset + env: + batch_size: 256 + replay: + odrl: true + +rl_dataloader: + num_workers: 2 + batch_size: null + + +trainer: + env_freq: null + rm_steps: 500000 + rl_steps: 500000 + log_freq: 500 + profile_freq: 500 + eval_freq: 10000 + label_reward: true + normalize_reward: true + +rm_eval: + function: eval_reward_model + eval_dataset_kwargs: + class: RPLComparisonDataset + env: + batch_size: 32 + label_key: rl_reward_sum + eval: true + replay: + odrl: true + +rl_eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 500000 + +processor: null diff --git a/scripts/configs/rpl/rpl_iql/default.yaml b/scripts/configs/rpl/rpl_iql/default.yaml new file mode 100644 index 0000000..c0cb5fe --- /dev/null +++ b/scripts/configs/rpl/rpl_iql/default.yaml @@ -0,0 +1,132 @@ +num_tasks: 4 +task_name: 0.1_0.5_2.0_5.0 + +algorithm: + class: RPL_IQL + num_tasks: # number of tasks + alpha: 0.7 # expectile for Eq. 8 + expectile: 0.7 # expectile of IQL + beta: 0.3333 # inv. temperature of IQL + max_exp_clip: 100.0 + rm_label: true + +checkpoint: null +seed: 0 +name: default +debug: false +device: null +wandb: + activate: false + entity: null + project: null + +# gravity: 1.0 # perform RL on gravity=1.0 +# variant: gravity-100 +# replay: false + +# env: HalfCheetah-v3 +# env_kwargs: +# env_wrapper: MujocoParamOverWrite +# env_wrapper_kwargs: +# overwrite_args: +# gravity: +# do_scale: True + +env: walker2d-gravity-1.0 +env_kwargs: +env_wrapper: +env_wrapper_kwargs: + +replay: true +label_key: rl_dir + +optim: + default: + class: Adam + lr: 0.0003 + +network: + actor: + class: SquashedGaussianActor + hidden_dims: [256, 256] + reparameterize: false + conditioned_logstd: false + logstd_min: -5 + logstd_max: 2 + critic: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256] + value: + class: Critic + ensemble_size: 1 + hidden_dims: [256, 256] + reward: + class: EnsembleMLP + ensemble_size: 1 + hidden_dims: [256, 256] + reward_act: identity #sigmoid + optimal: + class: Critic + ensemble_size: 2 + hidden_dims: [256, 256] + +rm_dataset: + - class: MultiRPLComparisonDataset # please fill in the class and parameters here + label_key: + num_tasks: + env: + task_name: + batch_size: 32 #8 + variant: gravity-100 + odrl: true + # capacity: 5000 +rm_dataloader: + num_workers: 2 + batch_size: null + +rl_dataset: + - class: RPLOfflineDataset + env: + batch_size: 256 + #variant: + replay: + odrl: true + +rl_dataloader: + num_workers: 2 + batch_size: null + + +trainer: + env_freq: null + rm_label: true + rm_steps: 1000000 + rl_steps: 500000 + log_freq: 500 + profile_freq: 500 + eval_freq: 5000 + +rm_eval: + function: eval_reward_model + eval_dataset_kwargs: + class: RPLComparisonDataset + env: + batch_size: 32 + label_key: rl_reward_sum # + #variant: + eval: true + replay: + odrl: true + +rl_eval: + function: eval_offline + num_ep: 10 + deterministic: true + +schedulers: + actor: + class: CosineAnnealingLR + T_max: 1000000 + +processor: null diff --git a/scripts/main.py b/scripts/main.py index dbaaf2d..95e9a67 100644 --- a/scripts/main.py +++ b/scripts/main.py @@ -12,7 +12,35 @@ from wiserl.trainer.offline_trainer import OfflineTrainer from wiserl.utils.utils import use_placeholder +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/halfcheetah-gravity-150.yaml --name halfcheetah-gravity-150-oracle-iql-rpl +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/halfcheetah-gravity-100.yaml --name halfcheetah-gravity-100-oracle-iql-rpl +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/halfcheetah-gravity-50.yaml --name halfcheetah-gravity-50-oracle-iql-rpl + +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/gravity-10.yaml --name Walker2d-v3-gravity-10-oracle-iql-rpl +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/gravity-50.yaml --name Walker2d-v3-gravity-50-oracle-iql-rpl +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/gravity-100.yaml --name Walker2d-v3-gravity-100-oracle-iql-rpl +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/gravity-150.yaml --name Walker2d-v3-gravity-150-oracle-iql-rpl + +# python scripts/main.py --config scripts/configs/oracle_awac/rpl/gravity-10.yaml --name Walker2d-v3-gravity-10-oracle-awac-rpl +# python scripts/main.py --config scripts/configs/oracle_awac/rpl/gravity-50.yaml --name Walker2d-v3-gravity-50-oracle-awac-rpl +# python scripts/main.py --config scripts/configs/oracle_awac/rpl/gravity-100.yaml --name Walker2d-v3-gravity-100-oracle-awac-rpl +# python scripts/main.py --config scripts/configs/oracle_awac/rpl/gravity-150.yaml --name Walker2d-v3-gravity-150-oracle-awac-rpl + +#python scripts/main.py --config scripts/configs/oracle_iql/rpl/gravity-100-test.yaml --name Walker2d-medium-gravity-100-test-oracle-iql + +# python scripts/main.py --config scripts/configs/oracle_awac/rpl/halfcheetah-gravity-150.yaml --name halfcheetah-gravity-150-oracle-awac-rpl +# python scripts/main.py --config scripts/configs/oracle_awac/rpl/halfcheetah-gravity-100.yaml --name halfcheetah-gravity-100-oracle-awac-rpl +# python scripts/main.py --config scripts/configs/oracle_awac/rpl/halfcheetah-gravity-50.yaml --name halfcheetah-gravity-50-oracle-awac-rpl + +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/halfcheetah-gravity-50-tr.yaml --name halfcheetah-gravity-50-oracle-awac-rpl-trajectory +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/halfcheetah-gravity-100-tr.yaml --name halfcheetah-gravity-100-oracle-awac-rpl-trajectory + + if __name__ == "__main__": + # import debugpy + # debugpy.listen(5678) + # debugpy.wait_for_client() + args = parse_args(convert=False, post_init=use_placeholder) name_prefix = f"{args['algorithm']['class']}/{args['name']}/{args['env']}" logger = CompositeLogger( @@ -28,10 +56,15 @@ ) logger.log_config(args, type="yaml") setup(args, logger) + import torch + args['device'] = torch.device('cuda:0') # process the environment env_fn = functools.partial(get_env, args["env"], args["env_kwargs"], args["env_wrapper"], args["env_wrapper_kwargs"]) - eval_env_fn = functools.partial(get_env, args["env"], args["env_kwargs"], args["env_wrapper"], args["env_wrapper_kwargs"]) + if "eval_env" in args: + eval_env_fn = functools.partial(get_env, args["eval_env"], args["eval_env_kwargs"], args["eval_env_wrapper"], args["eval_env_wrapper_kwargs"]) + else: + eval_env_fn = functools.partial(get_env, args["env"], args["env_kwargs"], args["env_wrapper"], args["env_wrapper_kwargs"]) env = env_fn() # define the algorithm diff --git a/scripts/main_online.py b/scripts/main_online.py new file mode 100644 index 0000000..01e489c --- /dev/null +++ b/scripts/main_online.py @@ -0,0 +1,93 @@ +import argparse +import functools +import os +import shutil + +from UtilsRL.exp import parse_args, setup +from UtilsRL.logger import CompositeLogger + +import wandb +import wiserl.algorithm +from wiserl.env import get_env +from wiserl.trainer.online_trainer import OnlineTrainer +from wiserl.utils.utils import use_placeholder + +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/halfcheetah-gravity-150.yaml --name halfcheetah-gravity-150-oracle-iql-rpl +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/halfcheetah-gravity-100.yaml --name halfcheetah-gravity-100-oracle-iql-rpl +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/halfcheetah-gravity-50.yaml --name halfcheetah-gravity-50-oracle-iql-rpl + +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/gravity-10.yaml --name Walker2d-v3-gravity-10-oracle-iql-rpl +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/gravity-50.yaml --name Walker2d-v3-gravity-50-oracle-iql-rpl +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/gravity-100.yaml --name Walker2d-v3-gravity-100-oracle-iql-rpl +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/gravity-150.yaml --name Walker2d-v3-gravity-150-oracle-iql-rpl + +# python scripts/main.py --config scripts/configs/oracle_awac/rpl/gravity-10.yaml --name Walker2d-v3-gravity-10-oracle-awac-rpl +# python scripts/main.py --config scripts/configs/oracle_awac/rpl/gravity-50.yaml --name Walker2d-v3-gravity-50-oracle-awac-rpl +# python scripts/main.py --config scripts/configs/oracle_awac/rpl/gravity-100.yaml --name Walker2d-v3-gravity-100-oracle-awac-rpl +# python scripts/main.py --config scripts/configs/oracle_awac/rpl/gravity-150.yaml --name Walker2d-v3-gravity-150-oracle-awac-rpl + +#python scripts/main.py --config scripts/configs/oracle_iql/rpl/gravity-100-test.yaml --name Walker2d-medium-gravity-100-test-oracle-iql + +# python scripts/main.py --config scripts/configs/oracle_awac/rpl/halfcheetah-gravity-150.yaml --name halfcheetah-gravity-150-oracle-awac-rpl +# python scripts/main.py --config scripts/configs/oracle_awac/rpl/halfcheetah-gravity-100.yaml --name halfcheetah-gravity-100-oracle-awac-rpl +# python scripts/main.py --config scripts/configs/oracle_awac/rpl/halfcheetah-gravity-50.yaml --name halfcheetah-gravity-50-oracle-awac-rpl + +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/halfcheetah-gravity-50-tr.yaml --name halfcheetah-gravity-50-oracle-awac-rpl-trajectory +# python scripts/main.py --config scripts/configs/oracle_iql/rpl/halfcheetah-gravity-100-tr.yaml --name halfcheetah-gravity-100-oracle-awac-rpl-trajectory + + +if __name__ == "__main__": + # import debugpy + # debugpy.listen(5678) + # debugpy.wait_for_client() + + args = parse_args(convert=False, post_init=use_placeholder) + name_prefix = f"{args['algorithm']['class']}/{args['name']}/{args['env']}" + logger = CompositeLogger( + log_dir=f"./log/{name_prefix}", + name="seed"+str(args["seed"]), + logger_config={ + "TensorboardLogger": {}, + "WandbLogger": {**args["wandb"], "config": args, "settings": wandb.Settings(_disable_stats=True)}, + "CsvLogger": {"activate": args.get("csv", False)} + }, + backup_stdout=True, + activate=not args["debug"] + ) + logger.log_config(args, type="yaml") + setup(args, logger) + + # process the environment + env_fn = functools.partial(get_env, args["env"], args["env_kwargs"], args["env_wrapper"], args["env_wrapper_kwargs"]) + if "eval_env" in args: + eval_env_fn = functools.partial(get_env, args["eval_env"], args["eval_env_kwargs"], args["eval_env_wrapper"], args["eval_env_wrapper_kwargs"]) + else: + eval_env_fn = functools.partial(get_env, args["env"], args["env_kwargs"], args["env_wrapper"], args["env_wrapper_kwargs"]) + env = env_fn() + + # define the algorithm + + algorithm = vars(wiserl.algorithm)[args["algorithm"].pop("class")]( + env.observation_space, + env.action_space, + args["network"], + args["optim"], + args["schedulers"], + args["processor"], + args["checkpoint"], + **args["algorithm"], + device=args["device"] + ) + + # define the trainer + trainer = OnlineTrainer( + algorithm=algorithm, + env_fn=env_fn, + eval_env_fn=eval_env_fn, + eval_kwargs=args["eval"], + **args["buffer"], + **args["trainer"], + logger=logger, + device=args["device"] + ) + trainer.train() diff --git a/scripts/rmb_main.py b/scripts/rmb_main.py index a1f8068..dce208d 100644 --- a/scripts/rmb_main.py +++ b/scripts/rmb_main.py @@ -12,7 +12,34 @@ from wiserl.trainer.rmb_offline_trainer import RewardModelBasedOfflineTrainer from wiserl.utils.utils import use_placeholder +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/halfcheetah-gravity-80-150.yaml --name halfcheetah-gravity-80-150-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/halfcheetah-gravity-80-80.yaml --name halfcheetah-gravity-80-80-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/halfcheetah-gravity-150-150.yaml --name halfcheetah-gravity-150-150-bt-iql-rpl + +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-50-50.yaml --name Walker2d-v3-gravity-50-50-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-100-100.yaml --name Walker2d-v3-gravity-100-100-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-150-150.yaml --name Walker2d-v3-gravity-150-150-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-50-100.yaml --name Walker2d-v3-gravity-50-100-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-50-150.yaml --name Walker2d-v3-gravity-50-150-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-100-150.yaml --name Walker2d-v3-gravity-100-150-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-150-50.yaml --name Walker2d-v3-gravity-150-50-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-150-100.yaml --name Walker2d-v3-gravity-150-100-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-100-100-b1-e9.yaml --name Walker2d-v3-gravity-100-100-bt-iql-rpl-b1-e9 + + +# python scripts/rmb_main.py --config scripts/configs/bt_awac/rpl/gravity-50-50.yaml --name Walker2d-v3-gravity-50-50-bt-awac-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_awac/rpl/gravity-100-100.yaml --name Walker2d-v3-gravity-100-100-bt-awac-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_awac/rpl/gravity-150-150.yaml --name Walker2d-v3-gravity-150-150-bt-awac-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_awac/rpl/gravity-50-100.yaml --name Walker2d-v3-gravity-50-100-bt-awac-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_awac/rpl/gravity-50-150.yaml --name Walker2d-v3-gravity-50-150-bt-awac-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_awac/rpl/gravity-150-50.yaml --name Walker2d-v3-gravity-150-50-bt-awac-rpl + + + if __name__ == "__main__": + # import debugpy + # debugpy.listen(5678) + # debugpy.wait_for_client() args = parse_args(convert=False, post_init=use_placeholder) name_prefix = f"{args['algorithm']['class']}/{args['name']}/{args['env']}" logger = CompositeLogger( @@ -28,10 +55,15 @@ ) logger.log_config(args, type="yaml") setup(args, logger) + import torch + args['device'] = torch.device('cuda:0') # process the environment env_fn = functools.partial(get_env, args["env"], args["env_kwargs"], args["env_wrapper"], args["env_wrapper_kwargs"]) - eval_env_fn = functools.partial(get_env, args["env"], args["env_kwargs"], args["env_wrapper"], args["env_wrapper_kwargs"]) + if "eval_env" in args: + eval_env_fn = functools.partial(get_env, args["eval_env"], args["eval_env_kwargs"], args["eval_env_wrapper"], args["eval_env_wrapper_kwargs"]) + else: + eval_env_fn = functools.partial(get_env, args["env"], args["env_kwargs"], args["env_wrapper"], args["env_wrapper_kwargs"]) env = env_fn() # define the algorithm diff --git a/scripts/rmb_main_online.py b/scripts/rmb_main_online.py new file mode 100644 index 0000000..4fee4ea --- /dev/null +++ b/scripts/rmb_main_online.py @@ -0,0 +1,95 @@ +import argparse +import functools +import os +import shutil + +from UtilsRL.exp import parse_args, setup +from UtilsRL.logger import CompositeLogger + +import wandb +import wiserl.algorithm +from wiserl.env import get_env +from wiserl.trainer.rmb_online_trainer import RewardModelBasedOnlineTrainer +from wiserl.utils.utils import use_placeholder + +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/halfcheetah-gravity-80-150.yaml --name halfcheetah-gravity-80-150-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/halfcheetah-gravity-80-80.yaml --name halfcheetah-gravity-80-80-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/halfcheetah-gravity-150-150.yaml --name halfcheetah-gravity-150-150-bt-iql-rpl + +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-50-50.yaml --name Walker2d-v3-gravity-50-50-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-100-100.yaml --name Walker2d-v3-gravity-100-100-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-150-150.yaml --name Walker2d-v3-gravity-150-150-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-50-100.yaml --name Walker2d-v3-gravity-50-100-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-50-150.yaml --name Walker2d-v3-gravity-50-150-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-100-150.yaml --name Walker2d-v3-gravity-100-150-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-150-50.yaml --name Walker2d-v3-gravity-150-50-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-150-100.yaml --name Walker2d-v3-gravity-150-100-bt-iql-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_iql/rpl/gravity-100-100-b1-e9.yaml --name Walker2d-v3-gravity-100-100-bt-iql-rpl-b1-e9 + + +# python scripts/rmb_main.py --config scripts/configs/bt_awac/rpl/gravity-50-50.yaml --name Walker2d-v3-gravity-50-50-bt-awac-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_awac/rpl/gravity-100-100.yaml --name Walker2d-v3-gravity-100-100-bt-awac-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_awac/rpl/gravity-150-150.yaml --name Walker2d-v3-gravity-150-150-bt-awac-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_awac/rpl/gravity-50-100.yaml --name Walker2d-v3-gravity-50-100-bt-awac-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_awac/rpl/gravity-50-150.yaml --name Walker2d-v3-gravity-50-150-bt-awac-rpl +# python scripts/rmb_main.py --config scripts/configs/bt_awac/rpl/gravity-150-50.yaml --name Walker2d-v3-gravity-150-50-bt-awac-rpl + +# + +if __name__ == "__main__": + # import debugpy + # debugpy.listen(5678) + # debugpy.wait_for_client() + args = parse_args(convert=False, post_init=use_placeholder) + name_prefix = f"{args['algorithm']['class']}/{args['name']}/{args['env']}" + logger = CompositeLogger( + log_dir=f"./log/{name_prefix}", + name="seed"+str(args["seed"]), + logger_config={ + "TensorboardLogger": {}, + "WandbLogger": {**args["wandb"], "config": args, "settings": wandb.Settings(_disable_stats=True)}, + "CsvLogger": {"activate": args.get("csv", False)} + }, + backup_stdout=True, + activate=not args["debug"] + ) + logger.log_config(args, type="yaml") + setup(args, logger) + + # process the environment + env_fn = functools.partial(get_env, args["env"], args["env_kwargs"], args["env_wrapper"], args["env_wrapper_kwargs"]) + if "eval_env" in args: + eval_env_fn = functools.partial(get_env, args["eval_env"], args["eval_env_kwargs"], args["eval_env_wrapper"], args["eval_env_wrapper_kwargs"]) + else: + eval_env_fn = functools.partial(get_env, args["env"], args["env_kwargs"], args["env_wrapper"], args["env_wrapper_kwargs"]) + env = env_fn() + + # define the algorithm + + algorithm = vars(wiserl.algorithm)[args["algorithm"].pop("class")]( + env.observation_space, + env.action_space, + args["network"], + args["optim"], + args["schedulers"], + args["processor"], + args["checkpoint"], + **args["algorithm"], + device=args["device"] + ) + + # define the trainer + trainer = RewardModelBasedOnlineTrainer( + algorithm=algorithm, + env_fn=env_fn, + eval_env_fn=eval_env_fn, + rm_dataset_kwargs=args["rm_dataset"], + rm_dataloader_kwargs=args["rm_dataloader"], + rm_eval_kwargs=args["rm_eval"], + **args["buffer"], + rl_eval_kwargs=args["rl_eval"], + **args["trainer"], + logger=logger, + device=args["device"] + ) + trainer.train() diff --git a/scripts/run_main.sh b/scripts/run_main.sh new file mode 100755 index 0000000..0720d00 --- /dev/null +++ b/scripts/run_main.sh @@ -0,0 +1,28 @@ +envs=("HalfCheetah-v3") +gravity_variant_pairs=( + "0.5:gravity-50" + "1.0:gravity-100" + "1.5:gravity-150" +) +algorithm=("oracle_awac") +info=${1:-""} + +# 遍历参数组合 +for env in "${envs[@]}"; do + for pair in "${gravity_variant_pairs[@]}"; do + # 分割 gravity 和 variant + gravity=$(echo $pair | cut -d':' -f1) + variant=$(echo $pair | cut -d':' -f2) + + # 输出当前组合 + echo "Running with env=$env, gravity=$gravity, variant=$variant" + + # 运行 Python 程序 + python scripts/main.py --config scripts/configs/${algorithm}/rpl/rpl.yaml \ + --name $env-$variant-${algorithm}-rpl-$info\ + --env $env\ + --gravity $gravity \ + --variant $variant + + done +done \ No newline at end of file diff --git a/scripts/run_main_replay.sh b/scripts/run_main_replay.sh new file mode 100755 index 0000000..a7959eb --- /dev/null +++ b/scripts/run_main_replay.sh @@ -0,0 +1,28 @@ +envs=("HalfCheetah-v3") +gravity_variant_pairs=( + "0.5:gravity-50" + "1.0:gravity-100" + "1.5:gravity-150" +) +algorithm=("oracle_awac") +info=${1:-""} +# 遍历参数组合 +for env in "${envs[@]}"; do + for pair in "${gravity_variant_pairs[@]}"; do + # 分割 gravity 和 variant + gravity=$(echo $pair | cut -d':' -f1) + variant=$(echo $pair | cut -d':' -f2) + + # 输出当前组合 + echo "Running with env=$env, gravity=$gravity, variant=$variant" + + # 运行 Python 程序 + python scripts/main.py --config scripts/configs/${algorithm}/rpl/rpl.yaml \ + --name $env-$variant-${algorithm}-rpl-replay-$info\ + --env $env\ + --gravity $gravity \ + --variant $variant\ + --dataset.0.replay true + + done +done \ No newline at end of file diff --git a/scripts/run_rmb_main.sh b/scripts/run_rmb_main.sh new file mode 100755 index 0000000..8548868 --- /dev/null +++ b/scripts/run_rmb_main.sh @@ -0,0 +1,39 @@ +envs=("HalfCheetah-v3") +src_gravity_variant_pairs=( + "0.5:gravity-50" + # "1.0:gravity-100" + # "1.5:gravity-150" +) +tgt_gravity_variant_pairs=( + # "0.5:gravity-50" + "1.0:gravity-100" + "1.5:gravity-150" +) +algorithm=("bt_awac") + +info=${1:-""} + +# 遍历参数组合 +for env in "${envs[@]}"; do + for pair1 in "${src_gravity_variant_pairs[@]}"; do + for pair2 in "${tgt_gravity_variant_pairs[@]}"; do + # 分割 gravity 和 variant + src_gravity=$(echo $pair1 | cut -d':' -f1) + src_variant=$(echo $pair1 | cut -d':' -f2) + tgt_gravity=$(echo $pair2 | cut -d':' -f1) + tgt_variant=$(echo $pair2 | cut -d':' -f2) + + # 输出当前组合 + echo "Running with env=$env, src_gravity=$src_gravity, tgt_gravity=$tgt_gravity" + + # 运行 Python 程序 + python scripts/rmb_main.py --config scripts/configs/${algorithm}/rpl/rpl.yaml \ + --name $env-$src_variant-$tgt_variant-${algorithm}-rpl$info \ + --env $env\ + --src_gravity $src_gravity \ + --src_variant $src_variant \ + --tgt_gravity $tgt_gravity \ + --tgt_variant $tgt_variant + done + done +done \ No newline at end of file diff --git a/scripts/run_rmb_main_replay.sh b/scripts/run_rmb_main_replay.sh new file mode 100755 index 0000000..a23e447 --- /dev/null +++ b/scripts/run_rmb_main_replay.sh @@ -0,0 +1,43 @@ +envs=("HalfCheetah-v3") +src_gravity_variant_pairs=( + "0.1:gravity-10" + "0.5:gravity-50" + "1.0:gravity-100" + "1.5:gravity-150" + "3.0:gravity-300" +) +tgt_gravity_variant_pairs=( + # "0.1:gravity-10" + # "0.5:gravity-50" + # "1.0:gravity-100" + # "1.5:gravity-150" + "3.0:gravity-300" +) +algorithm=("bt_awac") +info=${1:-""} + +# 遍历参数组合 +for env in "${envs[@]}"; do + for pair2 in "${tgt_gravity_variant_pairs[@]}"; do + for pair1 in "${src_gravity_variant_pairs[@]}"; do + # 分割 gravity 和 variant + src_gravity=$(echo $pair1 | cut -d':' -f1) + src_variant=$(echo $pair1 | cut -d':' -f2) + tgt_gravity=$(echo $pair2 | cut -d':' -f1) + tgt_variant=$(echo $pair2 | cut -d':' -f2) + + # 输出当前组合 + echo "Running with env=$env, src_gravity=$src_gravity, tgt_gravity=$tgt_gravity" + + # 运行 Python 程序 + python scripts/rmb_main.py --config scripts/configs/${algorithm}/rpl/rpl.yaml \ + --name $env-$src_variant-$tgt_variant-${algorithm}-rpl-replay$info \ + --env $env\ + --src_gravity $src_gravity \ + --src_variant $src_variant \ + --tgt_gravity $tgt_gravity \ + --tgt_variant $tgt_variant \ + --replay true + done + done +done \ No newline at end of file diff --git a/wiserl/algorithm/__init__.py b/wiserl/algorithm/__init__.py index 3d67fd4..d6aab91 100644 --- a/wiserl/algorithm/__init__.py +++ b/wiserl/algorithm/__init__.py @@ -12,3 +12,9 @@ from wiserl.algorithm.pt.pt_awac import PTAWAC from wiserl.algorithm.pt.pt_iql import PTIQL from wiserl.algorithm.sft import SFT +from wiserl.algorithm.oracle_td3bc import OracleTD3BC +from wiserl.algorithm.bt.bt_td3bc import BTTD3BC +from wiserl.algorithm.oracle_sac import OracleSAC +from wiserl.algorithm.bt.bt_sac import BTSAC +from wiserl.algorithm.rpl.rpl_iql import RPL_IQL +from wiserl.algorithm.rpl.rpl_awac import RPL_AWAC \ No newline at end of file diff --git a/wiserl/algorithm/bt/bt_awac.py b/wiserl/algorithm/bt/bt_awac.py index 38060a6..ab27ed0 100644 --- a/wiserl/algorithm/bt/bt_awac.py +++ b/wiserl/algorithm/bt/bt_awac.py @@ -115,7 +115,7 @@ def train_step(self, batches, step: int, total_steps: int) -> Dict: reward = self.select_reward({"obs": obs, "action": action}, deterministic=True) # compute the loss for actor - actor_loss, advantage = self.actor_loss(obs, action) + actor_loss, advantage, exp_advantage= self.actor_loss(obs, action) self.optim["actor"].zero_grad() actor_loss.backward() self.optim["actor"].step() @@ -136,7 +136,9 @@ def train_step(self, batches, step: int, total_steps: int) -> Dict: "loss/q_loss": q_loss.item(), "loss/actor_loss": actor_loss.item(), "misc/q_pred": q_pred.mean().item(), - "misc/advantage": advantage.mean().item() + "misc/advantage": advantage.mean().item(), + "misc/exp_advantage_mean": exp_advantage.mean().item(), + "misc/exp_advantage_std": exp_advantage.std().item() } return metrics diff --git a/wiserl/algorithm/bt/bt_iql.py b/wiserl/algorithm/bt/bt_iql.py index 17e8c52..66c3970 100644 --- a/wiserl/algorithm/bt/bt_iql.py +++ b/wiserl/algorithm/bt/bt_iql.py @@ -123,6 +123,7 @@ def train_step(self, batches, step: int, total_steps: int) -> Dict: self.target_network.eval() q_old = self.target_network.critic(obs, action) q_old = torch.min(q_old, dim=0)[0] + # q_old = torch.mean(q_old, dim=0)[0] # compute the loss for value network v_loss, v_pred = self.v_loss(obs.detach(), q_old) diff --git a/wiserl/algorithm/bt/bt_sac.py b/wiserl/algorithm/bt/bt_sac.py new file mode 100644 index 0000000..8e42fe4 --- /dev/null +++ b/wiserl/algorithm/bt/bt_sac.py @@ -0,0 +1,209 @@ +import itertools +import os +from operator import itemgetter +from typing import Any, Dict, Optional, Type + +import torch +import torch.nn as nn + +import wiserl.module +from wiserl.algorithm.oracle_sac import OracleSAC +from wiserl.utils.misc import make_target, sync_target + + +class BTSAC(OracleSAC): + def __init__( + self, + *args, + alpha: float = 0.2, + auto_alpha: bool = False, + discount: float = 0.99, + tau: float = 0.005, + reward_reg: float = 0.0, + rm_label: bool = True, + **kwargs + ) -> None: + super().__init__( + *args, + alpha=alpha, + auto_alpha=auto_alpha, + discount=discount, + tau=tau, + **kwargs + ) + self.reward_reg = reward_reg + self.rm_label = rm_label + self.obs_dim = self.observation_space.shape[0] + self.action_dim = self.action_space.shape[0] + + self.reward_criterion = torch.nn.BCEWithLogitsLoss(reduction="none") + + def setup_network(self, network_kwargs): + super().setup_network(network_kwargs) + reward_act = { + "identity": nn.Identity(), + "sigmoid": nn.Sigmoid(), + }.get(network_kwargs["reward"].pop("reward_act")) + reward = vars(wiserl.module)[network_kwargs["reward"].pop("class")]( + input_dim=self.observation_space.shape[0]+self.action_space.shape[0], + output_dim=1, + **network_kwargs["reward"] + ) + self.network["reward"] = nn.Sequential(self.network["encoder"], reward, reward_act) + + + def setup_optimizers(self, optim_kwargs): + super().setup_optimizers(optim_kwargs) + default_kwargs = optim_kwargs.get("default", {}) + reward_kwargs = default_kwargs.copy() + reward_kwargs.update(optim_kwargs.get("reward", {})) + self.optim["reward"] = vars(torch.optim)[reward_kwargs.pop("class")](self.network.reward.parameters(), **reward_kwargs) + + def select_action(self, batch, deterministic: bool=True): + return super().select_action(batch, deterministic) + + def select_reward(self, batch, deterministic=False): + obs, action = batch["obs"], batch["action"] + reward = self.network.reward(torch.concat([obs, action], dim=-1)) + return reward.mean(0).detach() + + def pretrain_step(self, batches, step: int, total_steps: int) -> Dict: + batch = batches[0] + F_B, F_S = batch["obs_1"].shape[0:2] + all_obs = torch.concat([ + batch["obs_1"].reshape(-1, self.obs_dim), + batch["obs_2"].reshape(-1, self.obs_dim) + ]) + all_action = torch.concat([ + batch["action_1"].reshape(-1, self.action_dim), + batch["action_2"].reshape(-1, self.action_dim) + ]) + self.network.reward.train() + all_reward = self.network.reward(torch.concat([all_obs, all_action], dim=-1)) + r1, r2 = torch.chunk(all_reward, 2, dim=1) + E = r1.shape[0] + r1, r2 = r1.reshape(E, F_B, F_S, 1), r2.reshape(E, F_B, F_S, 1) + logits = r2.sum(dim=2) - r1.sum(dim=2) + labels = batch["label"].float().unsqueeze(0).expand_as(logits) + reward_loss = self.reward_criterion(logits, labels).sum(0).mean() + reg_loss = (r1**2).sum(0).mean() + (r2**2).sum(0).mean() + with torch.no_grad(): + reward_accuracy = ((logits > 0) == torch.round(labels)).float().mean() + + self.optim["reward"].zero_grad() + (reward_loss + self.reward_reg * reg_loss).backward() + self.optim["reward"].step() + + metrics = { + "loss/reward_loss": reward_loss.item(), + "loss/reward_reg_loss": reg_loss.item(), + "misc/reward_acc": reward_accuracy.item(), + "misc/reward_value": all_reward.mean().item() + } + return metrics + + def train_step(self, batches, step: int, total_steps: int) -> Dict: + # rl_batch = batches[0] + # metrics = {} + # obs, action, next_obs, terminal = itemgetter("obs", "action", "next_obs", "terminal")(rl_batch) + # terminal = terminal.float() + # if self.rm_label: + # reward = itemgetter("reward")(rl_batch) + # else: + # with torch.no_grad(): + # reward = self.select_reward({"obs": obs, "action": action}, deterministic=True) + + # q_loss, q_metrics = self.q_loss(obs, action, next_obs, reward, terminal) + # metrics.update(q_metrics) + # self.optim["critic"].zero_grad() + # q_loss.backward() + # self.optim["critic"].step() + + # if step % self.actor_update_interval == 0: + # actor_loss, actor_metrics = self.actor_loss(obs, action) + # metrics.update(actor_metrics) + # self.optim["actor"].zero_grad() + # actor_loss.backward() + # self.optim["actor"].step() + + # sync_target(self.network.critic, self.target_network.critic, tau=self.tau) + # sync_target(self.network.actor, self.target_network.actor, tau=self.tau) + + # for _, scheduler in self.schedulers.items(): + # scheduler.step() + + # return metrics + if isinstance(batches, list): + batch, *_ = batches + elif isinstance(batches, dict): + batch = batches + else: + assert 0,f'Undefined Type:{type(batches)}' + metrics = {} + if "obs_1" in batch: + obs = torch.cat([batch["obs_1"], batch["obs_2"]], dim=0) # (B, S+1) + action = torch.cat([batch["action_1"], batch["action_2"]], dim=0) # (B, S+1) + reward = torch.cat([batch["reward_1"], batch["reward_2"]], dim=0) + terminal = torch.cat([batch["terminal_1"], batch["terminal_2"]], dim=0) + + encoded_obs = self.network.encoder(obs) + + q_loss, q_pred = self.q_loss( + encoded_obs[:, :-1].detach(), + action[:, :-1], + encoded_obs[:, 1:].detach(), + reward[:, :-1], + terminal[:, :-1] + ) + else: + obs = batch["obs"] + action = batch["action"] + reward = batch["reward"] + terminal = batch["terminal"].float() + next_obs = batch["next_obs"] + + encoded_obs = self.network.encoder(obs) + next_encoded_obs = self.network.encoder(next_obs) + + q_loss, q_metrics = self.q_loss(encoded_obs, action, next_encoded_obs, reward, terminal) + metrics.update(q_metrics) + self.optim["critic"].zero_grad() + q_loss.backward() + self.optim["critic"].step() + + # compute the loss for actor + actor_loss, actor_metrics= self.actor_loss(encoded_obs, action) + self.optim["actor"].zero_grad() + actor_loss.backward() + self.optim["actor"].step() + metrics.update(actor_metrics) + + if self._is_auto_alpha: + alpha_loss = self.alpha_loss(encoded_obs) + self.alpha_optim.zero_grad() + alpha_loss.backward() + self.alpha_optim.step() + self._alpha = self._log_alpha.exp().detach() + else: + alpha_loss = 0 + metrics["misc/alpha"] = self._alpha.item() + + if step % self.target_freq == 0: + sync_target(self.network.critic, self.target_network.critic, tau=self.tau) + + for _, scheduler in self.schedulers.items(): + scheduler.step() + + return metrics + + + def load_pretrain(self, path): + for attr in ["reward"]: + state_dict = torch.load(os.path.join(path, attr+".pt"), map_location=self.device) + self.network.__getattr__(attr).load_state_dict(state_dict) + + def save_pretrain(self, path): + os.makedirs(path, exist_ok=True) + for attr in ["reward"]: + state_dict = self.network.__getattr__(attr).state_dict() + torch.save(state_dict, os.path.join(path, attr+".pt")) diff --git a/wiserl/algorithm/bt/bt_td3bc.py b/wiserl/algorithm/bt/bt_td3bc.py new file mode 100644 index 0000000..6d7ef1c --- /dev/null +++ b/wiserl/algorithm/bt/bt_td3bc.py @@ -0,0 +1,201 @@ +import itertools +import os +from operator import itemgetter +from typing import Any, Dict, Optional, Type + +import torch +import torch.nn as nn + +import wiserl.module +from wiserl.algorithm.oracle_td3bc import OracleTD3BC +from wiserl.utils.misc import make_target, sync_target + + +class BTTD3BC(OracleTD3BC): + def __init__( + self, + *args, + alpha: float = 0.2, + policy_noise: float = 0.2, + noise_clip: float = 0.5, + max_action: float = 1.0, + discount: float = 0.99, + tau: float = 0.005, + actor_update_interval: int = 2, + reward_reg: float = 0.0, + rm_label: bool = True, + **kwargs + ) -> None: + super().__init__( + *args, + alpha=alpha, + policy_noise=policy_noise, + noise_clip=noise_clip, + max_action=max_action, + discount=discount, + tau=tau, + actor_update_interval=actor_update_interval, + **kwargs + ) + self.reward_reg = reward_reg + self.rm_label = rm_label + self.obs_dim = self.observation_space.shape[0] + self.action_dim = self.action_space.shape[0] + + self.reward_criterion = torch.nn.BCEWithLogitsLoss(reduction="none") + + def setup_network(self, network_kwargs): + super().setup_network(network_kwargs) + reward_act = { + "identity": nn.Identity(), + "sigmoid": nn.Sigmoid(), + }.get(network_kwargs["reward"].pop("reward_act")) + reward = vars(wiserl.module)[network_kwargs["reward"].pop("class")]( + input_dim=self.observation_space.shape[0]+self.action_space.shape[0], + output_dim=1, + **network_kwargs["reward"] + ) + self.network["reward"] = nn.Sequential(self.network["encoder"], reward, reward_act) + + + def setup_optimizers(self, optim_kwargs): + super().setup_optimizers(optim_kwargs) + default_kwargs = optim_kwargs.get("default", {}) + reward_kwargs = default_kwargs.copy() + reward_kwargs.update(optim_kwargs.get("reward", {})) + self.optim["reward"] = vars(torch.optim)[reward_kwargs.pop("class")](self.network.reward.parameters(), **reward_kwargs) + + def select_action(self, batch, deterministic: bool=True): + return super().select_action(batch, deterministic) + + def select_reward(self, batch, deterministic=False): + obs, action = batch["obs"], batch["action"] + reward = self.network.reward(torch.concat([obs, action], dim=-1)) + return reward.mean(0).detach() + + def pretrain_step(self, batches, step: int, total_steps: int) -> Dict: + batch = batches[0] + F_B, F_S = batch["obs_1"].shape[0:2] + all_obs = torch.concat([ + batch["obs_1"].reshape(-1, self.obs_dim), + batch["obs_2"].reshape(-1, self.obs_dim) + ]) + all_action = torch.concat([ + batch["action_1"].reshape(-1, self.action_dim), + batch["action_2"].reshape(-1, self.action_dim) + ]) + self.network.reward.train() + all_reward = self.network.reward(torch.concat([all_obs, all_action], dim=-1)) + r1, r2 = torch.chunk(all_reward, 2, dim=1) + E = r1.shape[0] + r1, r2 = r1.reshape(E, F_B, F_S, 1), r2.reshape(E, F_B, F_S, 1) + logits = r2.sum(dim=2) - r1.sum(dim=2) + labels = batch["label"].float().unsqueeze(0).expand_as(logits) + reward_loss = self.reward_criterion(logits, labels).sum(0).mean() + reg_loss = (r1**2).sum(0).mean() + (r2**2).sum(0).mean() + with torch.no_grad(): + reward_accuracy = ((logits > 0) == torch.round(labels)).float().mean() + + self.optim["reward"].zero_grad() + (reward_loss + self.reward_reg * reg_loss).backward() + self.optim["reward"].step() + + metrics = { + "loss/reward_loss": reward_loss.item(), + "loss/reward_reg_loss": reg_loss.item(), + "misc/reward_acc": reward_accuracy.item(), + "misc/reward_value": all_reward.mean().item() + } + return metrics + + def train_step(self, batches, step: int, total_steps: int) -> Dict: + # rl_batch = batches[0] + # metrics = {} + # obs, action, next_obs, terminal = itemgetter("obs", "action", "next_obs", "terminal")(rl_batch) + # terminal = terminal.float() + # if self.rm_label: + # reward = itemgetter("reward")(rl_batch) + # else: + # with torch.no_grad(): + # reward = self.select_reward({"obs": obs, "action": action}, deterministic=True) + + # q_loss, q_metrics = self.q_loss(obs, action, next_obs, reward, terminal) + # metrics.update(q_metrics) + # self.optim["critic"].zero_grad() + # q_loss.backward() + # self.optim["critic"].step() + + # if step % self.actor_update_interval == 0: + # actor_loss, actor_metrics = self.actor_loss(obs, action) + # metrics.update(actor_metrics) + # self.optim["actor"].zero_grad() + # actor_loss.backward() + # self.optim["actor"].step() + + # sync_target(self.network.critic, self.target_network.critic, tau=self.tau) + # sync_target(self.network.actor, self.target_network.actor, tau=self.tau) + + # for _, scheduler in self.schedulers.items(): + # scheduler.step() + + # return metrics + batch, *_ = batches + metrics = {} + if "obs_1" in batch: + obs = torch.cat([batch["obs_1"], batch["obs_2"]], dim=0) # (B, S+1) + action = torch.cat([batch["action_1"], batch["action_2"]], dim=0) # (B, S+1) + reward = torch.cat([batch["reward_1"], batch["reward_2"]], dim=0) + terminal = torch.cat([batch["terminal_1"], batch["terminal_2"]], dim=0) + + encoded_obs = self.network.encoder(obs) + + q_loss, q_pred = self.q_loss( + encoded_obs[:, :-1].detach(), + action[:, :-1], + encoded_obs[:, 1:].detach(), + reward[:, :-1], + terminal[:, :-1] + ) + else: + obs = batch["obs"] + action = batch["action"] + reward = batch["reward"] + terminal = batch["terminal"].float() + next_obs = batch["next_obs"] + + encoded_obs = self.network.encoder(obs) + next_encoded_obs = self.network.encoder(next_obs) + + q_loss, q_metrics = self.q_loss(encoded_obs, action, next_encoded_obs, reward, terminal) + metrics.update(q_metrics) + self.optim["critic"].zero_grad() + q_loss.backward() + self.optim["critic"].step() + + # compute the loss for actor + if step % self.actor_update_interval == 0: + actor_loss, actor_metrics = self.actor_loss(encoded_obs, action) + metrics.update(actor_metrics) + self.optim["actor"].zero_grad() + actor_loss.backward() + self.optim["actor"].step() + + sync_target(self.network.critic, self.target_network.critic, tau=self.tau) + sync_target(self.network.actor, self.target_network.actor, tau=self.tau) + + for _, scheduler in self.schedulers.items(): + scheduler.step() + + return metrics + + + def load_pretrain(self, path): + for attr in ["reward"]: + state_dict = torch.load(os.path.join(path, attr+".pt"), map_location=self.device) + self.network.__getattr__(attr).load_state_dict(state_dict) + + def save_pretrain(self, path): + os.makedirs(path, exist_ok=True) + for attr in ["reward"]: + state_dict = self.network.__getattr__(attr).state_dict() + torch.save(state_dict, os.path.join(path, attr+".pt")) diff --git a/wiserl/algorithm/oracle_awac.py b/wiserl/algorithm/oracle_awac.py index b9a0e8e..030f4be 100644 --- a/wiserl/algorithm/oracle_awac.py +++ b/wiserl/algorithm/oracle_awac.py @@ -83,7 +83,7 @@ def actor_loss(self, encoded_obs, action, reduce=True): elif isinstance(self.network.actor, GaussianActor): policy_out = - self.network.actor.evaluate(encoded_obs, action)[0] actor_loss = (exp_advantage * policy_out) - return actor_loss.mean() if reduce else actor_loss, advantage + return actor_loss.mean() if reduce else actor_loss, advantage, exp_advantage def q_loss(self, encoded_obs, action, next_encoded_obs, reward, terminal, reduce=True): with torch.no_grad(): @@ -97,26 +97,39 @@ def q_loss(self, encoded_obs, action, next_encoded_obs, reward, terminal, reduce def train_step(self, batches, step:int, total_steps: int): batch, *_ = batches - obs = torch.cat([batch["obs_1"], batch["obs_2"]], dim=0) # (B, S+1) - action = torch.cat([batch["action_1"], batch["action_2"]], dim=0) # (B, S+1) - reward = torch.cat([batch["reward_1"], batch["reward_2"]], dim=0) - terminal = torch.cat([batch["terminal_1"], batch["terminal_2"]], dim=0) - - encoded_obs = self.network.encoder(obs) - - q_loss, q_pred = self.q_loss( - encoded_obs[:, :-1].detach(), - action[:, :-1], - encoded_obs[:, 1:].detach(), - reward[:, :-1], - terminal[:, :-1] - ) + if "obs_1" in batch: + obs = torch.cat([batch["obs_1"], batch["obs_2"]], dim=0) # (B, S+1) + action = torch.cat([batch["action_1"], batch["action_2"]], dim=0) # (B, S+1) + reward = torch.cat([batch["reward_1"], batch["reward_2"]], dim=0) + terminal = torch.cat([batch["terminal_1"], batch["terminal_2"]], dim=0) + + encoded_obs = self.network.encoder(obs) + + q_loss, q_pred = self.q_loss( + encoded_obs[:, :-1].detach(), + action[:, :-1], + encoded_obs[:, 1:].detach(), + reward[:, :-1], + terminal[:, :-1] + ) + else: + obs = batch["obs"] + action = batch["action"] + reward = batch["reward"] + terminal = batch["terminal"].float() + next_obs = batch["next_obs"] + + encoded_obs = self.network.encoder(obs) + next_encoded_obs = self.network.encoder(next_obs) + + q_loss, q_pred = self.q_loss(encoded_obs, action, next_encoded_obs, reward, terminal) + self.optim["critic"].zero_grad() q_loss.backward() self.optim["critic"].step() # compute the loss for actor - actor_loss, advantage = self.actor_loss(encoded_obs, action) + actor_loss, advantage, exp_advantage= self.actor_loss(encoded_obs, action) self.optim["actor"].zero_grad() actor_loss.backward() self.optim["actor"].step() @@ -131,6 +144,8 @@ def train_step(self, batches, step:int, total_steps: int): "loss/q_loss": q_loss.item(), "loss/actor_loss": actor_loss.item(), "misc/q_pred": q_pred.mean().item(), - "misc/advantage": advantage.mean().item() + "misc/advantage": advantage.mean().item(), + "misc/exp_advantage_mean": exp_advantage.mean().item(), + "misc/exp_advantage_std": exp_advantage.std().item() } return metrics diff --git a/wiserl/algorithm/oracle_iql.py b/wiserl/algorithm/oracle_iql.py index 27eb81c..29a09ca 100644 --- a/wiserl/algorithm/oracle_iql.py +++ b/wiserl/algorithm/oracle_iql.py @@ -97,6 +97,7 @@ def actor_loss(self, encoded_obs, action, q_old, v, reduce=True): elif isinstance(self.network.actor, GaussianActor): policy_out = - self.network.actor.evaluate(encoded_obs, action)[0] actor_loss = (exp_advantage * policy_out) + # actor_loss = policy_out return actor_loss.mean() if reduce else actor_loss, advantage def q_loss(self, encoded_obs, action, next_encoded_obs, reward, terminal, reduce=True): @@ -109,10 +110,16 @@ def q_loss(self, encoded_obs, action, next_encoded_obs, reward, terminal, reduce def train_step(self, batches, step:int, total_steps: int): batch, *_ = batches - obs = torch.cat([batch["obs_1"], batch["obs_2"]], dim=0) # (B, S+1) - action = torch.cat([batch["action_1"], batch["action_2"]], dim=0) # (B, S+1) - reward = torch.cat([batch["reward_1"], batch["reward_2"]], dim=0) - terminal = torch.cat([batch["terminal_1"], batch["terminal_2"]], dim=0) + if "obs_1" in batch: + obs = torch.cat([batch["obs_1"], batch["obs_2"]], dim=0) # (B, S+1) + action = torch.cat([batch["action_1"], batch["action_2"]], dim=0) # (B, S+1) + reward = torch.cat([batch["reward_1"], batch["reward_2"]], dim=0) + terminal = torch.cat([batch["terminal_1"], batch["terminal_2"]], dim=0) + else: + obs = batch["obs"] + action = batch["action"] + reward = batch["reward"] + terminal = batch["terminal"].float() encoded_obs = self.network.encoder(obs) @@ -120,6 +127,7 @@ def train_step(self, batches, step:int, total_steps: int): self.target_network.eval() q_old = self.target_network.critic(encoded_obs, action) q_old = torch.min(q_old, dim=0)[0] + # q_old = torch.mean(q_old, dim=0)[0] # compute the loss for value network v_loss, v_pred = self.v_loss(encoded_obs.detach(), q_old) @@ -134,13 +142,20 @@ def train_step(self, batches, step:int, total_steps: int): self.optim["actor"].step() # compute the loss for q, offset by 1 - q_loss, q_pred = self.q_loss( - encoded_obs[:, :-1].detach(), - action[:, :-1], - encoded_obs[:, 1:].detach(), - reward[:, :-1], - terminal[:, :-1] - ) + # Using trajectory + # q_loss, q_pred = self.q_loss( + # encoded_obs[:, :-1].detach(), + # action[:, :-1], + # encoded_obs[:, 1:].detach(), + # reward[:, :-1], + # terminal[:, :-1] + # ) + + # Using transition + next_obs = batch["next_obs"] + next_encoded_obs = self.network.encoder(next_obs) + q_loss, q_pred = self.q_loss(encoded_obs, action, next_encoded_obs, reward, terminal) + self.optim["critic"].zero_grad() q_loss.backward() self.optim["critic"].step() diff --git a/wiserl/algorithm/oracle_sac.py b/wiserl/algorithm/oracle_sac.py new file mode 100644 index 0000000..0c7af7a --- /dev/null +++ b/wiserl/algorithm/oracle_sac.py @@ -0,0 +1,180 @@ +import itertools +from typing import Any, Dict, Optional, Type, Union, Tuple + +import torch +import torch.nn as nn + +import wiserl.module +from wiserl.algorithm.base import Algorithm +from wiserl.module.actor import DeterministicActor, GaussianActor +from wiserl.utils.misc import make_target, sync_target + + +class OracleSAC(Algorithm): + def __init__( + self, + *args, + alpha: Union[float, Tuple[float, float]] = 0.2, + auto_alpha: bool = False, + discount: float = 0.99, + tau: float = 0.005, + target_freq: int = 1, + **kwargs + ) -> None: + super().__init__(*args, **kwargs) + self._is_auto_alpha = auto_alpha + if self._is_auto_alpha: + target_entropy = -float(self.action_space.shape[-1]) + alpha_lr = alpha + self._log_alpha = nn.Parameter(torch.tensor([0.0], dtype=torch.float32, device=self.device), requires_grad=True) + self._target_entropy = target_entropy + self.alpha_optim = torch.optim.Adam([self._log_alpha], lr=alpha_lr) + self._alpha = self._log_alpha.detach().exp() + else: + self._alpha = torch.tensor([alpha], dtype=torch.float32, device=self.device, requires_grad=False) + + + self.target_freq = target_freq + self.discount = discount + self.tau = tau + + def setup_network(self, network_kwargs): + network = {} + network["actor"] = vars(wiserl.module)[network_kwargs["actor"].pop("class")]( + input_dim=self.observation_space.shape[0], + output_dim=self.action_space.shape[0], + **network_kwargs["actor"] + ) + network["critic"] = vars(wiserl.module)[network_kwargs["critic"].pop("class")]( + input_dim=self.observation_space.shape[0]+self.action_space.shape[0], + output_dim=1, + **network_kwargs["critic"] + ) + if "encoder" in network_kwargs: + network["encoder"] = vars(wiserl.module)[network_kwargs["encoder"].pop("class")]( + input_dim=self.observation_space.shape[0], + output_dim=1, + **network_kwargs["encoder"] + ) + else: + network["encoder"] = nn.Identity() + self.network = nn.ModuleDict(network) + self.target_network = nn.ModuleDict({ + "critic": make_target(self.network.critic) + }) + + def setup_optimizers(self, optim_kwargs): + self.optim = {} + default_kwargs = optim_kwargs.get("default", {}) + + actor_kwargs = default_kwargs.copy() + actor_kwargs.update(optim_kwargs.get("actor", {})) + actor_params = itertools.chain(self.network.actor.parameters(), self.network.encoder.parameters()) + self.optim["actor"] = vars(torch.optim)[actor_kwargs.pop("class")](actor_params, **actor_kwargs) + + critic_kwargs = default_kwargs.copy() + critic_kwargs.update(optim_kwargs.get("critic", {})) + self.optim["critic"] = vars(torch.optim)[critic_kwargs.pop("class")](self.network.critic.parameters(), **critic_kwargs) + + def select_action(self, batch, deterministic: bool=True): + obs = self.network.encoder(batch["obs"]) + action, *_ = self.network.actor.sample(obs, deterministic=deterministic) + return action.squeeze().cpu().numpy() + + def actor_loss(self, encoded_obs, action, reduce=True): + new_actions, new_logprobs, _ = self.network.actor.sample(encoded_obs) + q_values = self.network.critic(encoded_obs, new_actions) + if len(q_values.shape) == 2: + q_values = q_values.unsqueeze(0) + q_values_min = torch.min(q_values, dim=0)[0] + q_values_std = torch.std(q_values, dim=0).mean().item() + q_values_mean = q_values.mean().item() + actor_loss = self._alpha * new_logprobs - q_values_min + + return actor_loss.mean() if reduce else actor_loss, { + "loss/actor_loss": actor_loss.mean() if reduce else actor_loss,\ + "misc/q_values_std": q_values_std,\ + "misc/q_values_min": q_values_min.mean().item(),\ + "misc/q_values_mean": q_values_mean} + + def q_loss(self, encoded_obs, action, next_encoded_obs, reward, terminal, reduce=True): + with torch.no_grad(): + self.target_network.eval() + next_actions, next_logprobs, _ = self.network.actor.sample(next_encoded_obs) + target_q = self.target_network.critic(next_encoded_obs, next_actions).min(0)[0]- self._alpha * next_logprobs + target_q = reward + self.discount * (1-terminal) * target_q + q_pred = self.network.critic(encoded_obs, action) + q_loss = (q_pred - target_q.unsqueeze(0)).pow(2).sum(0) + return q_loss.mean() if reduce else q_loss, {"loss/q_loss": q_loss.mean() if reduce else q_loss, "misc/q_pred":q_pred.mean() if reduce else q_pred} + + def alpha_loss(self, encoded_obs, reduce=True): + with torch.no_grad(): + _, new_logprobs, _ = self.network.actor.sample(encoded_obs) + alpha_loss = -(self._log_alpha * (new_logprobs + self._target_entropy)).mean() + return alpha_loss.mean() if reduce else alpha_loss + + def train_step(self, batches, step:int, total_steps: int): + if isinstance(batches, list): + batch, *_ = batches + elif isinstance(batches, dict): + batch = batches + else: + assert 0,f'Undefined Type:{type(batches)}' + metrics = {} + if "obs_1" in batch: + obs = torch.cat([batch["obs_1"], batch["obs_2"]], dim=0) # (B, S+1) + action = torch.cat([batch["action_1"], batch["action_2"]], dim=0) # (B, S+1) + reward = torch.cat([batch["reward_1"], batch["reward_2"]], dim=0) + terminal = torch.cat([batch["terminal_1"], batch["terminal_2"]], dim=0) + + encoded_obs = self.network.encoder(obs) + + q_loss, q_pred = self.q_loss( + encoded_obs[:, :-1].detach(), + action[:, :-1], + encoded_obs[:, 1:].detach(), + reward[:, :-1], + terminal[:, :-1] + ) + else: + obs = batch["obs"] + action = batch["action"] + reward = batch["reward"] + terminal = batch["terminal"].float() + next_obs = batch["next_obs"] + + encoded_obs = self.network.encoder(obs) + next_encoded_obs = self.network.encoder(next_obs) + + q_loss, q_metrics = self.q_loss(encoded_obs, action, next_encoded_obs, reward, terminal) + + metrics.update(q_metrics) + self.optim["critic"].zero_grad() + q_loss.backward() + self.optim["critic"].step() + + # compute the loss for actor + actor_loss, actor_metrics= self.actor_loss(encoded_obs, action) + self.optim["actor"].zero_grad() + actor_loss.backward() + self.optim["actor"].step() + metrics.update(actor_metrics) + + if self._is_auto_alpha: + alpha_loss = self.alpha_loss(encoded_obs) + self.alpha_optim.zero_grad() + alpha_loss.backward() + self.alpha_optim.step() + self._alpha = self._log_alpha.exp().detach() + else: + alpha_loss = 0 + metrics["misc/alpha"] = self._alpha.item() + + if step % self.target_freq == 0: + sync_target(self.network.critic, self.target_network.critic, tau=self.tau) + + for _, scheduler in self.schedulers.items(): + scheduler.step() + + + return metrics diff --git a/wiserl/algorithm/oracle_td3bc.py b/wiserl/algorithm/oracle_td3bc.py new file mode 100644 index 0000000..0da101b --- /dev/null +++ b/wiserl/algorithm/oracle_td3bc.py @@ -0,0 +1,162 @@ +import itertools +from typing import Any, Dict, Optional, Type + +import torch +import torch.nn as nn + +import wiserl.module +from wiserl.algorithm.base import Algorithm +from wiserl.utils.misc import make_target, sync_target + + +class OracleTD3BC(Algorithm): + def __init__( + self, + *args, + alpha: float = 0.2, + policy_noise: float = 0.2, + noise_clip: float = 0.5, + max_action: float = 1.0, + discount: float = 0.99, + tau: float = 0.005, + actor_update_interval: int = 2, + **kwargs + ) -> None: + super().__init__(*args, **kwargs) + self.alpha = alpha + self.policy_noise = policy_noise + self.noise_clip = noise_clip + self.max_action = max_action + self.actor_update_interval = actor_update_interval + self.discount = discount + self.tau = tau + + def setup_network(self, network_kwargs): + network = {} + network["actor"] = vars(wiserl.module)[network_kwargs["actor"].pop("class")]( + input_dim=self.observation_space.shape[0], + output_dim=self.action_space.shape[0], + **network_kwargs["actor"] + ) + network["critic"] = vars(wiserl.module)[network_kwargs["critic"].pop("class")]( + input_dim=self.observation_space.shape[0]+self.action_space.shape[0], + output_dim=1, + **network_kwargs["critic"] + ) + if "encoder" in network_kwargs: + network["encoder"] = vars(wiserl.module)[network_kwargs["encoder"].pop("class")]( + input_dim=self.observation_space.shape[0], + output_dim=1, + **network_kwargs["encoder"] + ) + else: + network["encoder"] = nn.Identity() + self.network = nn.ModuleDict(network) + self.target_network = nn.ModuleDict({ + "actor": make_target(self.network.actor), + "critic": make_target(self.network.critic) + }) + + def setup_optimizers(self, optim_kwargs): + self.optim = {} + default_kwargs = optim_kwargs.get("default", {}) + + actor_kwargs = default_kwargs.copy() + actor_kwargs.update(optim_kwargs.get("actor", {})) + actor_params = itertools.chain(self.network.actor.parameters(), self.network.encoder.parameters()) + self.optim["actor"] = vars(torch.optim)[actor_kwargs.pop("class")](actor_params, **actor_kwargs) + + critic_kwargs = default_kwargs.copy() + critic_kwargs.update(optim_kwargs.get("critic", {})) + self.optim["critic"] = vars(torch.optim)[critic_kwargs.pop("class")](self.network.critic.parameters(), **critic_kwargs) + + def select_action(self, batch, deterministic: bool=True): + obs = self.network.encoder(batch["obs"]) + action, *_ = self.network.actor.sample(obs, deterministic=deterministic) + return action.squeeze().cpu().numpy() + + def actor_loss(self, encoded_obs, action, reduce=True): + new_actions = self.network.actor.sample(encoded_obs)[0] + new_q1 = self.network.critic(encoded_obs, new_actions)[0, ...] + bc_loss = torch.nn.functional.mse_loss(new_actions, action, reduce=None) + # print(bc_loss.requires_grad) + # bc_loss = torch.sum((new_actions - action)**2, dim=-1, keepdim=True) + # print(bc_loss.requires_grad) + # from wiserl.module.actor import DeterministicActor, GaussianActor + # if isinstance(self.network.actor, DeterministicActor): + # bc_loss = torch.sum((self.network.actor.sample(encoded_obs)[0] - action)**2, dim=-1, keepdim=True) + # elif isinstance(self.network.actor, GaussianActor): + # bc_loss = - self.network.actor.evaluate(encoded_obs, action)[0] + q_loss = - self.alpha / (new_q1.abs().mean().detach()) * new_q1 + total_loss = bc_loss + q_loss + #print(total_loss.requires_grad) + #assert 0 + return total_loss.mean() if reduce else total_loss, { + "loss/q_guide_loss": q_loss.mean().item(), + "loss/bc_loss": bc_loss.mean().item(), + } + + def q_loss(self, encoded_obs, action, next_encoded_obs, reward, terminal, reduce=True): + with torch.no_grad(): + self.target_network.eval() + next_actions = self.target_network.actor.sample(next_encoded_obs)[0] + noise = (torch.randn_like(next_actions) * self.policy_noise).clip(-self.noise_clip, self.noise_clip) + next_actions = (next_actions+noise).clip(-self.max_action, self.max_action) + target_q = self.target_network.critic(next_encoded_obs, next_actions).min(0)[0] + target_q = reward + self.discount * (1-terminal) * target_q + q_pred = self.network.critic(encoded_obs, action) + q_loss = (q_pred - target_q.unsqueeze(0)).pow(2).sum(0) + return q_loss.mean() if reduce else q_loss, { + "loss/critic_loss": q_loss.mean().item(), + "misc/q_pred": q_pred.mean().item(), + } + + def train_step(self, batches, step:int, total_steps: int): + batch, *_ = batches + metrics = {} + if "obs_1" in batch: + obs = torch.cat([batch["obs_1"], batch["obs_2"]], dim=0) # (B, S+1) + action = torch.cat([batch["action_1"], batch["action_2"]], dim=0) # (B, S+1) + reward = torch.cat([batch["reward_1"], batch["reward_2"]], dim=0) + terminal = torch.cat([batch["terminal_1"], batch["terminal_2"]], dim=0) + + encoded_obs = self.network.encoder(obs) + + q_loss, q_pred = self.q_loss( + encoded_obs[:, :-1].detach(), + action[:, :-1], + encoded_obs[:, 1:].detach(), + reward[:, :-1], + terminal[:, :-1] + ) + else: + obs = batch["obs"] + action = batch["action"] + reward = batch["reward"] + terminal = batch["terminal"].float() + next_obs = batch["next_obs"] + + encoded_obs = self.network.encoder(obs) + next_encoded_obs = self.network.encoder(next_obs) + + q_loss, q_metrics = self.q_loss(encoded_obs, action, next_encoded_obs, reward, terminal) + metrics.update(q_metrics) + self.optim["critic"].zero_grad() + q_loss.backward() + self.optim["critic"].step() + + # compute the loss for actor + if step % self.actor_update_interval == 0: + actor_loss, actor_metrics = self.actor_loss(encoded_obs, action) + metrics.update(actor_metrics) + self.optim["actor"].zero_grad() + actor_loss.backward() + self.optim["actor"].step() + + sync_target(self.network.critic, self.target_network.critic, tau=self.tau) + sync_target(self.network.actor, self.target_network.actor, tau=self.tau) + + for _, scheduler in self.schedulers.items(): + scheduler.step() + + return metrics diff --git a/wiserl/algorithm/rpl/rpl_awac.py b/wiserl/algorithm/rpl/rpl_awac.py new file mode 100644 index 0000000..b5338da --- /dev/null +++ b/wiserl/algorithm/rpl/rpl_awac.py @@ -0,0 +1,212 @@ +import itertools +import os +from operator import itemgetter +from typing import Any, Dict, Optional, Type + +import numpy as np +import torch +import torch.nn as nn + +import wiserl.module +from wiserl.algorithm.oracle_awac import OracleAWAC +from wiserl.module.actor import DeterministicActor, GaussianActor +from wiserl.utils.functional import expectile_regression +from wiserl.utils.misc import make_target, sync_target + + +class RPL_AWAC(OracleAWAC): + def __init__( + self, + *args, + num_tasks: int = 4, + alpha: float = 0.7, + beta: float = 0.3333, + reward_reg: float = 0.0, + max_exp_clip: float = 100.0, + discount: float = 0.99, + tau: float = 0.005, + target_freq: int = 1, + rm_label: bool = True, + **kwargs + ): + self.num_tasks = num_tasks + super().__init__( + *args, + beta=beta, + max_exp_clip=max_exp_clip, + discount=discount, + tau=tau, + target_freq=target_freq, + **kwargs + ) + + self.rm_label = rm_label + self.alpha = alpha + self.reward_reg = reward_reg + self.reward_criterion = torch.nn.BCEWithLogitsLoss(reduction="none") + self.obs_dim = self.observation_space.shape[0] + self.action_dim = self.action_space.shape[0] + + def setup_network(self, network_kwargs): + super().setup_network(network_kwargs) + reward_act = { + "identity": nn.Identity(), + "sigmoid": nn.Sigmoid(), + }.get(network_kwargs["reward"].pop("reward_act")) + reward = vars(wiserl.module)[network_kwargs["reward"].pop("class")]( + input_dim=self.observation_space.shape[0]+self.action_space.shape[0], + output_dim=1, + **network_kwargs["reward"] + ) + optimal = vars(wiserl.module)[network_kwargs["optimal"].pop("class")]( + input_dim=self.observation_space.shape[0], + output_dim=1, + ensemble_size=self.num_tasks*network_kwargs["optimal"]["ensemble_size"] + ) + self.network["reward"] = nn.Sequential(self.network["encoder"], reward, reward_act) + self.network["optimal"] = optimal + self.target_network["reward"] = make_target(self.network["reward"]) + self.target_network["optimal"] = make_target(self.network["optimal"]) + + def setup_optimizers(self, optim_kwargs): + super().setup_optimizers(optim_kwargs) + default_kwargs = optim_kwargs.get("default", {}) + reward_kwargs = default_kwargs.copy() + reward_kwargs.update(optim_kwargs.get("reward", {})) + optimal_kwargs = default_kwargs.copy() + optimal_kwargs.update(optim_kwargs.get("optimal", {})) + self.optim["reward"] = vars(torch.optim)[reward_kwargs.pop("class")](self.network.reward.parameters(), **reward_kwargs) + self.optim["optimal"] = vars(torch.optim)[optimal_kwargs.pop("class")](self.network.optimal.parameters(), **optimal_kwargs) + + def select_reward(self, batch, deterministic=False): + obs, action = batch["obs"], batch["action"] + reward = self.network.reward(torch.concat([obs, action], dim=-1)) + return reward.mean(0).detach() + + def get_optimal_values(self, model, obs, task_id): + original_shape = obs.shape[:-1] + optimal_values = model(obs).reshape(self.num_tasks, -1, *original_shape, 1) + optimal_values = torch.gather( + optimal_values, + 0, + task_id.unsqueeze(0).unsqueeze(-1).expand(*optimal_values.shape[1:]).unsqueeze(0).to(torch.int64) + ) + return optimal_values[0] # squeeze out the first dim + + def pretrain_step(self, batches, step: int, total_steps: int) -> Dict: + batch = batches[0] + B, S = batch["obs_1"].shape[0:2] + all_obs = torch.concat([ + batch["obs_1"], + batch["obs_2"] + ], dim=0) + all_action = torch.concat([ + batch["action_1"], + batch["action_2"] + ], dim=0) + all_next_obs = torch.concat([ + batch["next_obs_1"], + batch["next_obs_2"] + ], dim=0) + all_task_id = torch.concat([ + batch["task_id"], + batch["task_id"] + ], dim=0) + + # train the optimal networks + with torch.no_grad(): + all_reward_target = self.target_network.reward(torch.concat([all_obs, all_action], dim=-1))[0] # (B, L, 1) + all_optimal_target = self.get_optimal_values( + self.target_network.optimal, + all_next_obs, + all_task_id + ) # (E, B, L, 1) + all_target = all_optimal_target.min(0)[0] + all_target = all_reward_target + self.discount * all_target # CHECK: default no terminal + all_optimal_pred = self.get_optimal_values( + self.network.optimal, + all_obs, + all_task_id + ) # (E, B, L, 1) + optimal_loss = expectile_regression(all_optimal_pred, all_target.unsqueeze(0), expectile=self.alpha) + optimal_loss = optimal_loss.sum(0).mean() + + self.optim["optimal"].zero_grad() + optimal_loss.backward() + self.optim["optimal"].step() + + # train the reward networks + all_reward = self.network.reward(torch.concat([all_obs, all_action], dim=-1))[0] # (B, L, 1) + + all_adv = all_reward + self.discount * all_target - all_optimal_pred.mean(0).detach() + adv1, adv2 = torch.chunk(all_adv, 2, dim=0) + logits = adv2.sum(dim=1) - adv1.sum(dim=1) + labels = batch["label"].float() + reward_loss = self.reward_criterion(logits, labels).mean() + reg_loss = (all_reward**2).mean() + with torch.no_grad(): + reward_accuracy = ((logits > 0) == torch.round(labels)).float().mean() + + self.optim["reward"].zero_grad() + (reward_loss + self.reward_reg * reg_loss).backward() + self.optim["reward"].step() + + sync_target(self.network.reward, self.target_network.reward, tau=self.tau) + sync_target(self.network.optimal, self.target_network.optimal, tau=self.tau) + + metrics = { + "loss/reward_loss": reward_loss.item(), + "loss/reward_reg_loss": reg_loss.item(), + "loss/optimal_loss": optimal_loss.item(), + "misc/reward_acc": reward_accuracy.item(), + "misc/reward_value": all_reward.mean().item(), + "misc/optimal_value": all_optimal_pred.mean().item() + } + return metrics + + def train_step(self, batches, step: int, total_steps: int) -> Dict: + rl_batch = batches[0] + obs, action, next_obs, terminal = itemgetter("obs", "action", "next_obs", "terminal")(rl_batch) + terminal = terminal.float() + if self.rm_label: + reward = itemgetter("reward")(rl_batch) + else: + with torch.no_grad(): + reward = self.select_reward({"obs": obs, "action": action}, deterministic=True) + + actor_loss, advantage, exp_advantage= self.actor_loss(obs, action) + self.optim["actor"].zero_grad() + actor_loss.backward() + self.optim["actor"].step() + + q_loss, q_pred = self.q_loss(obs, action, next_obs, reward, terminal) + self.optim["critic"].zero_grad() + q_loss.backward() + self.optim["critic"].step() + + for _, scheduler in self.schedulers.items(): + scheduler.step() + + if step % self.target_freq == 0: + sync_target(self.network.critic, self.target_network.critic, tau=self.tau) + + metrics = { + "loss/q_loss": q_loss.item(), + "loss/actor_loss": actor_loss.item(), + "misc/q_pred": q_pred.mean().item(), + "misc/advantage": advantage.mean().item(), + "misc/exp_advantage_mean": exp_advantage.mean().item(), + "misc/exp_advantage_std": exp_advantage.std().item() + } + return metrics + + def load_pretrain(self, path): + for attr in ["reward"]: + state_dict = torch.load(os.path.join(path, attr+".pt"), map_location=self.device) + self.network.__getattr__(attr).load_state_dict(state_dict) + + def save_pretrain(self, path): + os.makedirs(path, exist_ok=True) + for attr in ["reward"]: + state_dict = self.network.__getattr__(attr).state_dict() + torch.save(state_dict, os.path.join(path, attr+".pt")) \ No newline at end of file diff --git a/wiserl/algorithm/rpl/rpl_iql.py b/wiserl/algorithm/rpl/rpl_iql.py index 019af3f..81a8c14 100644 --- a/wiserl/algorithm/rpl/rpl_iql.py +++ b/wiserl/algorithm/rpl/rpl_iql.py @@ -27,12 +27,11 @@ def __init__( discount: float = 0.99, tau: float = 0.005, target_freq: int = 1, + rm_label: bool = True, **kwargs - ): - self.num_tasks = num_tasks - self.alpha = alpha - self.reward_reg = reward_reg - self.reward_criterion = torch.nn.BCEWithLogitsLoss(reduction="none") + ): + self.num_tasks = num_tasks + # self.reward_criterion = torch.nn.BCEWithLogitsLoss(reduction="none") super().__init__( *args, expectile=expectile, @@ -43,6 +42,11 @@ def __init__( target_freq=target_freq, **kwargs ) + + self.rm_label = rm_label + self.alpha = alpha + self.reward_reg = reward_reg + self.reward_criterion = torch.nn.BCEWithLogitsLoss(reduction="none") self.obs_dim = self.observation_space.shape[0] self.action_dim = self.action_space.shape[0] @@ -53,14 +57,14 @@ def setup_network(self, network_kwargs): "sigmoid": nn.Sigmoid() }.get(network_kwargs["reward"].pop("reward_act")) reward = vars(wiserl.module)[network_kwargs["reward"].pop("class")]( - input_dim=self.observation_space.shape[0]+self.action_space.shape, + input_dim=self.observation_space.shape[0]+self.action_space.shape[0], output_dim=1, **network_kwargs["reward"] ) optimal = vars(wiserl.module)[network_kwargs["optimal"].pop("class")]( input_dim=self.observation_space.shape[0], output_dim=1, - ensemble_size=self.num_tasks*network_kwargs["opt"]["ensemble_size"] + ensemble_size=self.num_tasks*network_kwargs["optimal"]["ensemble_size"] ) self.network["reward"] = nn.Sequential(self.network["encoder"], reward, reward_act) @@ -89,9 +93,9 @@ def get_optimal_values(self, model, obs, task_id): optimal_values = torch.gather( optimal_values, 0, - task_id.unsqueeze(0).expand(*optimal_values.shape[1:]).unsqueeze(0) + task_id.unsqueeze(0).unsqueeze(-1).expand(*optimal_values.shape[1:]).unsqueeze(0).to(torch.int64) ) - return optimal_values + return optimal_values[0] # squeeze out the first dim def pretrain_step(self, batches, step: int, total_steps: int) -> Dict: batch = batches[0] @@ -115,20 +119,20 @@ def pretrain_step(self, batches, step: int, total_steps: int) -> Dict: # train the optimal networks with torch.no_grad(): - all_reward_target = self.target_network.reward(torch.concat([all_obs, all_action], dim=-1)) + all_reward_target = self.target_network.reward(torch.concat([all_obs, all_action], dim=-1))[0] # (B, L, 1) all_optimal_target = self.get_optimal_values( self.target_network.optimal, all_next_obs, all_task_id - ) + ) # (E, B, L, 1) all_target = all_optimal_target.min(0)[0] all_target = all_reward_target + self.discount * all_target # CHECK: default no terminal all_optimal_pred = self.get_optimal_values( self.network.optimal, all_obs, all_task_id - ) - optimal_loss = expectile_regression(all_optimal_pred.unsqueeze(0), all_target, expectile=self.expectile) + ) # (E, B, L, 1) + optimal_loss = expectile_regression(all_optimal_pred, all_target.unsqueeze(0), expectile=self.alpha) optimal_loss = optimal_loss.sum(0).mean() self.optim["optimal"].zero_grad() @@ -136,20 +140,24 @@ def pretrain_step(self, batches, step: int, total_steps: int) -> Dict: self.optim["optimal"].step() # train the reward networks - all_reward = self.network.reward(torch.concat([all_obs, all_action], dim=-1)) - all_adv = all_reward + self.discount * all_target - all_optimal_pred.detach() + all_reward = self.network.reward(torch.concat([all_obs, all_action], dim=-1))[0] # (B, L, 1) + + all_adv = all_reward + self.discount * all_target - all_optimal_pred.mean(0).detach() adv1, adv2 = torch.chunk(all_adv, 2, dim=0) logits = adv2.sum(dim=1) - adv1.sum(dim=1) labels = batch["label"].float() reward_loss = self.reward_criterion(logits, labels).mean() reg_loss = (all_reward**2).mean() with torch.no_grad(): - reward_accuracy = ((logits > 0) == torch.round(labels)).float() + reward_accuracy = ((logits > 0) == torch.round(labels)).float().mean() self.optim["reward"].zero_grad() (reward_loss + self.reward_reg * reg_loss).backward() self.optim["reward"].step() + sync_target(self.network.reward, self.target_network.reward, tau=self.tau) + sync_target(self.network.optimal, self.target_network.optimal, tau=self.tau) + metrics = { "loss/reward_loss": reward_loss.item(), "loss/reward_reg_loss": reg_loss.item(), @@ -174,6 +182,7 @@ def train_step(self, batches, step: int, total_steps: int) -> Dict: self.target_network.eval() q_old = self.target_network.critic(obs, action) q_old = torch.min(q_old, dim=0)[0] + # q_old = torch.mean(q_old, dim=0)[0] # compute the loss for value network v_loss, v_pred = self.v_loss(obs.detach(), q_old) diff --git a/wiserl/dataset/__init__.py b/wiserl/dataset/__init__.py index 195da3f..2446080 100644 --- a/wiserl/dataset/__init__.py +++ b/wiserl/dataset/__init__.py @@ -11,10 +11,16 @@ MetaworldComparisonDataset, MetaworldOfflineDataset, ) +from wiserl.dataset.rpl_dataset import ( + RPLComparisonDataset, + RPLOfflineDataset, +) from wiserl.dataset.mismatched_mujoco_dataset import ( MismatchedComparisonDataset, MismatchedOfflineDataset, ) +from wiserl.dataset.multi_rpl_dataset import MultiRPLComparisonDataset +from wiserl.dataset.apl_dataset import APLOfflineDataset from .replay_buffer import ReplayBuffer diff --git a/wiserl/dataset/apl_dataset.py b/wiserl/dataset/apl_dataset.py new file mode 100644 index 0000000..824a21b --- /dev/null +++ b/wiserl/dataset/apl_dataset.py @@ -0,0 +1,294 @@ +import os +import pickle +from typing import Dict, Optional, Union + +import d4rl +import gym +import numpy as np +import torch + +from wiserl.utils import utils + +prefix = "datasets/rpl" + +class APLComparisonDataset(torch.utils.data.IterableDataset): + def __init__( + self, + observation_space, + action_space, + env: str, + segment_length: Optional[int] = None, + batch_size: Optional[int] = None, + capacity: Optional[int] = None, + label_key: str="rl_sum", + odrl: bool = False, + variant: str = "gravity-50", + eval: bool = False, + replay: bool = False, + ): + super().__init__() + + self.env_name = env + self.batch_size = 1 if batch_size is None else batch_size + self.segment_length = segment_length + self.label_key = label_key + '_label' + self.variant = variant + self.eval = eval + train_or_eval = "eval" if eval else "train" + mid_name = f"collect_odrl/{self.env_name}" if odrl else f"{self.env_name}/{variant}" + if replay: + path = f"{prefix}/{mid_name}/replay_preference_{train_or_eval}_data.npz" + else: + path = f"{prefix}/{mid_name}/preference_{train_or_eval}_data.npz" + with open(path, "rb") as f: + data = np.load(f) + data = utils.nest_dict(data) + if capacity is not None: + data = utils.get_from_batch(data, 0, capacity) + data = utils.remove_float64(data) + lim = 1 - 1e-8 + data["action_1"] = np.clip(data["action_1"], a_min=-lim, a_max=lim) + data["action_2"] = np.clip(data["action_2"], a_min=-lim, a_max=lim) + + self.data = data + self.data_size, self.data_segment_length = data["action_1"].shape[:2] + + def __len__(self): + return self.data_size + + def sample_idx(self, idx): + idx = np.squeeze(idx) + is_batch = len(idx.shape) > 0 + if self.segment_length is not None: + start_idx = np.random.randint(self.data_segment_length - self.segment_length) + end_idx = start_idx + self.segment_length + else: + start_idx, end_idx = 0, self.data_segment_length + batch = { + "obs_1": self.data["obs_1"][idx, start_idx:end_idx], + "obs_2": self.data["obs_2"][idx, start_idx:end_idx], + "action_1": self.data["action_1"][idx, start_idx:end_idx], + "action_2": self.data["action_2"][idx, start_idx:end_idx], + "label": self.data[self.label_key][idx][:, None], + "reward_1": self.data["reward_1"][idx, start_idx:end_idx], + "reward_2": self.data["reward_2"][idx, start_idx:end_idx], + "terminal_1": np.zeros([len(idx), end_idx-start_idx, 1], dtype=np.float32) \ + if is_batch else np.zeros([end_idx-start_idx, 1], dtype=np.float32), + "terminal_2": np.zeros([len(idx), end_idx-start_idx, 1], dtype=np.float32) \ + if is_batch else np.zeros([end_idx-start_idx, 1], dtype=np.float32) + } + return batch + + def __iter__(self): + while True: + idxs = np.random.randint(0, len(self), size=self.batch_size) + yield self.sample_idx(idxs) + + def create_sequential_iter(self): + start, end = 0, min(self.batch_size, self.data_size) + while start < self.data_size: + idxs = list(range(start, min(end, self.data_size))) + yield self.sample_idx(idxs) + start += self.batch_size + end += self.batch_size + + +class APLOfflineDataset(torch.utils.data.IterableDataset): + def __init__( + self, + observation_space: gym.Space, + action_space: gym.Space, + env: str, + # segment_length: Optional[int] = None, + batch_size: Optional[int] = None, + capacity: Optional[int] = None, + mode: str = "transition", + odrl: bool = False, + variant: str = "gravity-50", + eval: bool = False, + replay: bool = True, + mismatch: str = "", + ): + super().__init__() + assert mode in {"transition", "trajectory"} + self.mode = mode + self.env_name = env + self.batch_size = 1 if batch_size is None else batch_size + # self.segment_length = segment_length + self.capacity = capacity + self.odrl = odrl + self.variant = variant + self.eval = eval + self.replay = replay + + self.mismatch = mismatch + self.load_dataset() + + def __len__(self): + return self.data_size + + def __iter__(self): + while True: + idxs = np.random.randint(0, self.data_size, self.batch_size) + idxs = np.squeeze(idxs) + traj_len = self.data["obs"][0].shape[0] + mask = np.ones([self.batch_size, traj_len, 1], dtype=np.float32) + timestep = np.stack([np.arange(traj_len) for _ in idxs], axis=0) + yield { + "obs": self.data["obs"][idxs], + "next_obs": self.data["next_obs"][idxs], + "action": self.data["action"][idxs], + "reward": self.data["reward"][idxs], + "terminal": self.data["terminal"][idxs], + "mask": mask, + "timestep": timestep, + } + + def load_dataset(self): + # Using preference datasets + mid_name = f"collect_odrl/{self.env_name}" if self.odrl else f"{self.env_name}/{self.variant}" + if len(self.mismatch)!=0: + tmp = self.env_name+'-to-'+self.mismatch + mid_name = f"collect_odrl_mismatch/{tmp}" + print("hh", tmp) + if self.mode == "trajectory": + train_or_eval = "eval" if self.eval else "train" + replay_or_none = 'replay_' if self.replay else "" + + path = f"{prefix}/{mid_name}/{replay_or_none}preference_{train_or_eval}_data.npz" + with open(path, "rb") as f: + data = np.load(f) + data = utils.nest_dict(data) + if self.capacity is not None: + data = utils.get_from_batch(data, 0, self.capacity) + data = utils.remove_float64(data) + lim = 1 - 1e-8 + data["action_1"] = np.clip(data["action_1"], a_min=-lim, a_max=lim) + data["action_2"] = np.clip(data["action_2"], a_min=-lim, a_max=lim) + N, L = data["obs_1"].shape[:2] + print(data.keys()) + + data = { + "obs": np.stack([data["obs_1"], data["obs_2"]], axis=0).reshape(2*N, L, -1), + "next_obs": np.stack([data["next_obs_1"], data["next_obs_2"]], axis=0).reshape(2*N, L, -1), + "action": np.stack([data["action_1"], data["action_2"]], axis=0).reshape(2*N, L, -1), + "reward": np.stack([data["reward_1"], data["reward_2"]], axis=0).reshape(2*N, L, -1), + "terminal": np.stack([data["terminal_1"], data["terminal_2"]], axis=0).reshape(2*N, L, -1), + } + + data = { + "obs": data["obs"], + "next_obs":data["next_obs"], + "action": data["action"], + "reward": data["reward"], + "terminal": data["terminal"], + } + data["mask"] = np.ones([2*N, L, 1], dtype=np.float32) + + self.traj_len = np.asarray([o.shape[0] for o in data["obs"]]) + self.data_size = len(self.traj_len) + else: + # Using offline datasets + if self.replay: + path = f"{prefix}/{mid_name}/labeled_replay.npz" + else: + path = f"{prefix}/{mid_name}/data.npz" + + with open(path, "rb") as f: + data = np.load(f) + data = utils.nest_dict(data) + if self.capacity is not None: + data = utils.get_from_batch(data, 0, self.capacity) + data = utils.remove_float64(data) + self.traj_len = np.sum(data['mask'],axis=-1) + obs_ = [] + next_obs_ = [] + action_ = [] + reward_ = [] + terminal_ = [] + timeout_ = [] + lim = 1 - 1e-8 + data["action"] = np.clip(data["action"], a_min=-lim, a_max=lim) + data['timeout'] = np.squeeze(data['timeout']) + + data['a1'] = data['q_value']-data['v_value'] + #data['a2'] = data['reward'] + data['next_v_value']-data['v_value'] + data['a2'] = data['reward'] + data['next_v_value']*(np.ones(data['terminal'].shape)-data['terminal'])-data['v_value'] + + for i in range(data['obs'].shape[0]): + obs_.extend(data['obs'][i][:int(self.traj_len[i])]) + next_obs_.extend(data['next_obs'][i][:int(self.traj_len[i])]) + action_.extend(data['action'][i][:int(self.traj_len[i])]) + + reward_.extend(data['reward'][i][:int(self.traj_len[i])]) + #reward_.extend(data['a1'][i][:int(self.traj_len[i])]) + #reward_.extend(data['a2'][i][:int(self.traj_len[i])]) + + terminal_.extend(data['terminal'][i][:int(self.traj_len[i])]) + timeout_.extend(data['timeout'][i][:int(self.traj_len[i])]) + + data = { + "obs": np.asarray(obs_), + "action": np.asarray(action_), + "next_obs": np.asarray(next_obs_), + "reward": np.asarray(reward_), + "terminal": np.asarray(terminal_), + "timeout": np.asarray(timeout_), + "mask": np.ones([len(obs_), 1], dtype=np.float32), + } + self.data_size = data["obs"].shape[0] + + if self.capacity is not None: + if self.capacity > self.data_size: + print(f"[Warning]: capacity {self.capacity} exceeds dataset size {self.data_size}") + self.data_size = min(self.data_size, self.capacity) + data = { + k: data[k][:self.data_size] for k in data + } + self.traj_len = self.traj_len[:self.data_size] + self.data = data + + @torch.no_grad() + def relabel_reward(self, agent): + assert hasattr(agent, "select_reward"), f"Agent {agent} must support relabel_reward!" + bs = 256 + for i_batch in range((self.data_size-1) // bs + 1): + idx = np.arange(i_batch*bs, min((i_batch+1)*bs, self.data_size)) + batch = { + "obs": self.data["obs"][idx], + "action": self.data["action"][idx], + "next_obs": self.data["next_obs"][idx], + "mask": self.data["mask"][idx] + } + batch = agent.format_batch(batch) + reward = agent.select_reward(batch).detach().cpu().numpy() + reward = reward * self.data["mask"][idx] + self.data["reward"][idx] = reward + + def normalize_reward(self): + if self.mode == "trajectory": + return_ = self.data["reward"].sum(1) + max_return = max( + abs(return_.max()), + abs(return_.min()), + return_.max() - return_.min(), + 1.0 + ) + # norm = 500. / max_return + # #print(f"norm: {norm} ") + # self.data["reward"] *= norm + # print(f"[RPLOfflineDataset]: return range: [{return_.min()}, {return_.max()}], multiplying norm factor {norm}.") + print(f"[APLOfflineDataset]: return range: [{return_.min()}, {return_.max()}].") + else: + ep_reward_ = [] + episode_reward = 0 + N = self.data["reward"].shape[0] + for i in range(N): + episode_reward += self.data["reward"][i] + if self.data["terminal"][i] or self.data["timeout"][i]: + ep_reward_.append(episode_reward) + episode_reward = 0 + max_return = max(abs(min(ep_reward_)).item(), abs(max(ep_reward_)).item(), (max(ep_reward_)-min(ep_reward_)).item(), 1.0) + norm = 1000 / max_return + self.data["reward"] *= norm + print(f"[APLOfflineDataset]: return range: [{min(ep_reward_)}, {max(ep_reward_)}], multiplying norm factor {norm}.") \ No newline at end of file diff --git a/wiserl/dataset/d4rl_dataset.py b/wiserl/dataset/d4rl_dataset.py index b676ab7..8eb80de 100644 --- a/wiserl/dataset/d4rl_dataset.py +++ b/wiserl/dataset/d4rl_dataset.py @@ -113,7 +113,11 @@ def __iter__(self): } def load_dataset(self): - env = gym.make(self.env_name) + # rewrite env name + env_name = self.env_name.split('-')[0] + env_name += '-medium' + env = gym.make(env_name) + #env = gym.make(self.env_name) dataset = env.get_dataset() N = dataset["rewards"].shape[0] @@ -247,7 +251,8 @@ def relabel_reward(self, agent): reward = agent.select_reward(batch).detach().cpu().numpy() reward = reward * self.data["mask"][idx] self.data["reward"][idx] = reward - + + def normalize_reward(self): if self.mode == "trajectory": # CHECK: may be bug. the max and min returns are not consistent with those computed in transition mode return_ = self.data["reward"].copy() diff --git a/wiserl/dataset/multi_rpl_dataset.py b/wiserl/dataset/multi_rpl_dataset.py new file mode 100644 index 0000000..6a27f24 --- /dev/null +++ b/wiserl/dataset/multi_rpl_dataset.py @@ -0,0 +1,109 @@ +import os +import pickle +from typing import Dict, Optional, Union + +import d4rl +import gym +import numpy as np +import torch + +from wiserl.utils import utils + +prefix = "datasets/rpl" + +class MultiRPLComparisonDataset(torch.utils.data.IterableDataset): + def __init__( + self, + observation_space, + action_space, + env: str, + num_tasks: int = 4, + task_name: str = '1.0_0.5_0.1_5.0', + segment_length: Optional[int] = None, + batch_size: Optional[int] = None, + capacity: Optional[int] = None, + label_key: str="rl_sum", + odrl: bool = False, + variant: str = "gravity-50", + eval: bool = False, + replay: bool = False, + ): + super().__init__() + + self.env_name = env + self.batch_size = 1 if batch_size is None else batch_size + self.segment_length = segment_length + self.label_key = label_key + '_label' + self.variant = variant + self.eval = eval + train_or_eval = "eval" if eval else "train" + #mid_name = f"collect_odrl/{self.env_name}" if odrl else f"{self.env_name}/{variant}" + self.tasks = task_name.split('_') + task_prefix = '-'.join(env.split('-')[:-1]) + #assert len(self.tasks) == num_tasks + num_tasks = len(self.tasks) + + for i in range(num_tasks): + mid_name = f"collect_odrl/{task_prefix+'-'+self.tasks[i]}" if odrl else f"{self.env_name}/{task_prefix+variant}" + if replay: + path = f"{prefix}/{mid_name}/replay_preference_{train_or_eval}_data.npz" + else: + path = f"{prefix}/{mid_name}/preference_{train_or_eval}_data.npz" + with open(path, "rb") as f: + data = np.load(f) + data = utils.nest_dict(data) + if capacity is not None: + data = utils.get_from_batch(data, 0, capacity) + data = utils.remove_float64(data) + lim = 1 - 1e-8 + data["action_1"] = np.clip(data["action_1"], a_min=-lim, a_max=lim) + data["action_2"] = np.clip(data["action_2"], a_min=-lim, a_max=lim) + data["task_id"] = np.full_like(data["reward_1"], i) + if i == 0: + self.data = data + else: + for key, value in data.items(): + self.data[key] = np.concatenate((self.data[key], value), axis=0) + self.data_size, self.data_segment_length = self.data["action_1"].shape[:2] + + def __len__(self): + return self.data_size + + def sample_idx(self, idx): + idx = np.squeeze(idx) + is_batch = len(idx.shape) > 0 + if self.segment_length is not None: + start_idx = np.random.randint(self.data_segment_length - self.segment_length) + end_idx = start_idx + self.segment_length + else: + start_idx, end_idx = 0, self.data_segment_length + batch = { + "obs_1": self.data["obs_1"][idx, start_idx:end_idx], + "obs_2": self.data["obs_2"][idx, start_idx:end_idx], + "next_obs_1": self.data["next_obs_1"][idx, start_idx:end_idx], + "next_obs_2": self.data["next_obs_2"][idx, start_idx:end_idx], + "action_1": self.data["action_1"][idx, start_idx:end_idx], + "action_2": self.data["action_2"][idx, start_idx:end_idx], + "label": self.data[self.label_key][idx][:, None], + "reward_1": self.data["reward_1"][idx, start_idx:end_idx], + "reward_2": self.data["reward_2"][idx, start_idx:end_idx], + "task_id": self.data["task_id"][idx, start_idx:end_idx], + "terminal_1": np.zeros([len(idx), end_idx-start_idx, 1], dtype=np.float32) \ + if is_batch else np.zeros([end_idx-start_idx, 1], dtype=np.float32), + "terminal_2": np.zeros([len(idx), end_idx-start_idx, 1], dtype=np.float32) \ + if is_batch else np.zeros([end_idx-start_idx, 1], dtype=np.float32) + } + return batch + + def __iter__(self): + while True: + idxs = np.random.randint(0, len(self), size=self.batch_size) + yield self.sample_idx(idxs) + + def create_sequential_iter(self): + start, end = 0, min(self.batch_size, self.data_size) + while start < self.data_size: + idxs = list(range(start, min(end, self.data_size))) + yield self.sample_idx(idxs) + start += self.batch_size + end += self.batch_size diff --git a/wiserl/dataset/rpl_dataset.py b/wiserl/dataset/rpl_dataset.py new file mode 100644 index 0000000..adaf487 --- /dev/null +++ b/wiserl/dataset/rpl_dataset.py @@ -0,0 +1,278 @@ +import os +import pickle +from typing import Dict, Optional, Union + +import d4rl +import gym +import numpy as np +import torch + +from wiserl.utils import utils + +prefix = "datasets/rpl" + +class RPLComparisonDataset(torch.utils.data.IterableDataset): + def __init__( + self, + observation_space, + action_space, + env: str, + segment_length: Optional[int] = None, + batch_size: Optional[int] = None, + capacity: Optional[int] = None, + label_key: str="rl_sum", + odrl: bool = False, + variant: str = "gravity-50", + eval: bool = False, + replay: bool = False, + ): + super().__init__() + + self.env_name = env + self.batch_size = 1 if batch_size is None else batch_size + self.segment_length = segment_length + self.label_key = label_key + '_label' + self.variant = variant + self.eval = eval + train_or_eval = "eval" if eval else "train" + mid_name = f"collect_odrl/{self.env_name}" if odrl else f"{self.env_name}/{variant}" + if replay: + path = f"{prefix}/{mid_name}/replay_preference_{train_or_eval}_data.npz" + else: + path = f"{prefix}/{mid_name}/preference_{train_or_eval}_data.npz" + with open(path, "rb") as f: + data = np.load(f) + data = utils.nest_dict(data) + if capacity is not None: + data = utils.get_from_batch(data, 0, capacity) + data = utils.remove_float64(data) + lim = 1 - 1e-8 + data["action_1"] = np.clip(data["action_1"], a_min=-lim, a_max=lim) + data["action_2"] = np.clip(data["action_2"], a_min=-lim, a_max=lim) + + self.data = data + self.data_size, self.data_segment_length = data["action_1"].shape[:2] + + def __len__(self): + return self.data_size + + def sample_idx(self, idx): + idx = np.squeeze(idx) + is_batch = len(idx.shape) > 0 + if self.segment_length is not None: + start_idx = np.random.randint(self.data_segment_length - self.segment_length) + end_idx = start_idx + self.segment_length + else: + start_idx, end_idx = 0, self.data_segment_length + batch = { + "obs_1": self.data["obs_1"][idx, start_idx:end_idx], + "obs_2": self.data["obs_2"][idx, start_idx:end_idx], + "action_1": self.data["action_1"][idx, start_idx:end_idx], + "action_2": self.data["action_2"][idx, start_idx:end_idx], + "label": self.data[self.label_key][idx][:, None], + "reward_1": self.data["reward_1"][idx, start_idx:end_idx], + "reward_2": self.data["reward_2"][idx, start_idx:end_idx], + "terminal_1": np.zeros([len(idx), end_idx-start_idx, 1], dtype=np.float32) \ + if is_batch else np.zeros([end_idx-start_idx, 1], dtype=np.float32), + "terminal_2": np.zeros([len(idx), end_idx-start_idx, 1], dtype=np.float32) \ + if is_batch else np.zeros([end_idx-start_idx, 1], dtype=np.float32) + } + return batch + + def __iter__(self): + while True: + idxs = np.random.randint(0, len(self), size=self.batch_size) + yield self.sample_idx(idxs) + + def create_sequential_iter(self): + start, end = 0, min(self.batch_size, self.data_size) + while start < self.data_size: + idxs = list(range(start, min(end, self.data_size))) + yield self.sample_idx(idxs) + start += self.batch_size + end += self.batch_size + + +class RPLOfflineDataset(torch.utils.data.IterableDataset): + def __init__( + self, + observation_space: gym.Space, + action_space: gym.Space, + env: str, + # segment_length: Optional[int] = None, + batch_size: Optional[int] = None, + capacity: Optional[int] = None, + mode: str = "transition", + odrl: bool = False, + variant: str = "gravity-50", + eval: bool = False, + replay: bool = False, + ): + super().__init__() + assert mode in {"transition", "trajectory"} + self.mode = mode + self.env_name = env + self.batch_size = 1 if batch_size is None else batch_size + # self.segment_length = segment_length + self.capacity = capacity + self.odrl = odrl + self.variant = variant + self.eval = eval + self.replay = replay + + self.load_dataset() + + def __len__(self): + return self.data_size + + def __iter__(self): + while True: + idxs = np.random.randint(0, self.data_size, self.batch_size) + idxs = np.squeeze(idxs) + traj_len = self.data["obs"][0].shape[0] + mask = np.ones([self.batch_size, traj_len, 1], dtype=np.float32) + timestep = np.stack([np.arange(traj_len) for _ in idxs], axis=0) + yield { + "obs": self.data["obs"][idxs], + "next_obs": self.data["next_obs"][idxs], + "action": self.data["action"][idxs], + "reward": self.data["reward"][idxs], + "terminal": self.data["terminal"][idxs], + "mask": mask, + "timestep": timestep, + } + + def load_dataset(self): + # Using preference datasets + mid_name = f"collect_odrl/{self.env_name}" if self.odrl else f"{self.env_name}/{self.variant}" + if self.mode == "trajectory": + train_or_eval = "eval" if self.eval else "train" + replay_or_none = 'replay_' if self.replay else "" + + path = f"{prefix}/{mid_name}/{replay_or_none}preference_{train_or_eval}_data.npz" + with open(path, "rb") as f: + data = np.load(f) + data = utils.nest_dict(data) + if self.capacity is not None: + data = utils.get_from_batch(data, 0, self.capacity) + data = utils.remove_float64(data) + lim = 1 - 1e-8 + data["action_1"] = np.clip(data["action_1"], a_min=-lim, a_max=lim) + data["action_2"] = np.clip(data["action_2"], a_min=-lim, a_max=lim) + N, L = data["obs_1"].shape[:2] + + data = { + "obs": np.stack([data["obs_1"], data["obs_2"]], axis=0).reshape(2*N, L, -1), + "next_obs": np.stack([data["next_obs_1"], data["next_obs_2"]], axis=0).reshape(2*N, L, -1), + "action": np.stack([data["action_1"], data["action_2"]], axis=0).reshape(2*N, L, -1), + "reward": np.stack([data["reward_1"], data["reward_2"]], axis=0).reshape(2*N, L, -1), + "terminal": np.stack([data["terminal_1"], data["terminal_2"]], axis=0).reshape(2*N, L, -1), + } + + data = { + "obs": data["obs"], + "next_obs":data["next_obs"], + "action": data["action"], + "reward": data["reward"], + "terminal": data["terminal"], + } + data["mask"] = np.ones([2*N, L, 1], dtype=np.float32) + + self.traj_len = np.asarray([o.shape[0] for o in data["obs"]]) + self.data_size = len(self.traj_len) + else: + # Using offline datasets + if self.replay: + path = f"{prefix}/{mid_name}/replay.npz" + else: + path = f"{prefix}/{mid_name}/data.npz" + + with open(path, "rb") as f: + data = np.load(f) + data = utils.nest_dict(data) + if self.capacity is not None: + data = utils.get_from_batch(data, 0, self.capacity) + data = utils.remove_float64(data) + self.traj_len = np.sum(data['mask'],axis=-1) + obs_ = [] + next_obs_ = [] + action_ = [] + reward_ = [] + terminal_ = [] + timeout_ = [] + lim = 1 - 1e-8 + data["action"] = np.clip(data["action"], a_min=-lim, a_max=lim) + data['timeout'] = np.squeeze(data['timeout']) + for i in range(data['obs'].shape[0]): + obs_.extend(data['obs'][i][:int(self.traj_len[i])]) + next_obs_.extend(data['next_obs'][i][:int(self.traj_len[i])]) + action_.extend(data['action'][i][:int(self.traj_len[i])]) + reward_.extend(data['reward'][i][:int(self.traj_len[i])]) + terminal_.extend(data['terminal'][i][:int(self.traj_len[i])]) + timeout_.extend(data['timeout'][i][:int(self.traj_len[i])]) + + data = { + "obs": np.asarray(obs_), + "action": np.asarray(action_), + "next_obs": np.asarray(next_obs_), + "reward": np.asarray(reward_), + "terminal": np.asarray(terminal_), + "timeout": np.asarray(timeout_), + "mask": np.ones([len(obs_), 1], dtype=np.float32), + } + self.data_size = data["obs"].shape[0] + + if self.capacity is not None: + if self.capacity > self.data_size: + print(f"[Warning]: capacity {self.capacity} exceeds dataset size {self.data_size}") + self.data_size = min(self.data_size, self.capacity) + data = { + k: data[k][:self.data_size] for k in data + } + self.traj_len = self.traj_len[:self.data_size] + self.data = data + + @torch.no_grad() + def relabel_reward(self, agent): + assert hasattr(agent, "select_reward"), f"Agent {agent} must support relabel_reward!" + bs = 256 + for i_batch in range((self.data_size-1) // bs + 1): + idx = np.arange(i_batch*bs, min((i_batch+1)*bs, self.data_size)) + batch = { + "obs": self.data["obs"][idx], + "action": self.data["action"][idx], + "next_obs": self.data["next_obs"][idx], + "mask": self.data["mask"][idx] + } + batch = agent.format_batch(batch) + reward = agent.select_reward(batch).detach().cpu().numpy() + reward = reward * self.data["mask"][idx] + self.data["reward"][idx] = reward + + def normalize_reward(self): + if self.mode == "trajectory": + return_ = self.data["reward"].sum(1) + max_return = max( + abs(return_.max()), + abs(return_.min()), + return_.max() - return_.min(), + 1.0 + ) + # norm = 500. / max_return + # #print(f"norm: {norm} ") + # self.data["reward"] *= norm + # print(f"[RPLOfflineDataset]: return range: [{return_.min()}, {return_.max()}], multiplying norm factor {norm}.") + print(f"[RPLOfflineDataset]: return range: [{return_.min()}, {return_.max()}].") + else: + ep_reward_ = [] + episode_reward = 0 + N = self.data["reward"].shape[0] + for i in range(N): + episode_reward += self.data["reward"][i] + if self.data["terminal"][i] or self.data["timeout"][i]: + ep_reward_.append(episode_reward) + episode_reward = 0 + max_return = max(abs(min(ep_reward_)).item(), abs(max(ep_reward_)).item(), (max(ep_reward_)-min(ep_reward_)).item(), 1.0) + norm = 1000 / max_return + self.data["reward"] *= norm + print(f"[D4RLOfflineDataset]: return range: [{min(ep_reward_)}, {max(ep_reward_)}], multiplying norm factor {norm}.") \ No newline at end of file diff --git a/wiserl/env/__init__.py b/wiserl/env/__init__.py index a5fa76d..0c7ca59 100644 --- a/wiserl/env/__init__.py +++ b/wiserl/env/__init__.py @@ -5,6 +5,7 @@ import gym import UtilsRL.env.wrapper from gym.envs import register +from wiserl.env.odrl_envs.mujoco.call_mujoco_env import call_mujoco_env from .base import EmptyEnv from .cliffwalking_env import CliffWalkingEnv @@ -53,7 +54,10 @@ def get_env( env_kwargs = env_kwargs or {} env = extra_envs[env](**env_kwargs) except KeyError as e: - env = gym.make(env, **env_kwargs) + try: + env = call_mujoco_env(env) + except Exception as e: + env = gym.make(env, **env_kwargs) if wrapper_class is not None: wrapper_kwargs = wrapper_kwargs or {} env = vars(UtilsRL.env.wrapper)[wrapper_class](env, **wrapper_kwargs) diff --git a/wiserl/env/odrl_envs/__init__.py b/wiserl/env/odrl_envs/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/wiserl/env/odrl_envs/adroit/__init__.py b/wiserl/env/odrl_envs/adroit/__init__.py new file mode 100644 index 0000000..ee07d79 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/__init__.py @@ -0,0 +1,189 @@ +from gym.envs.registration import register + +from pathlib import Path +import sys +sys.path.append(str(Path(__file__).parent.absolute())) + +# from door import DoorEnvV0 +# from relocate import RelocateEnvV0 +# from hammer import HammerEnvV0 +# from pen import PenEnvV0 +import door +import pen +import relocate +import hammer + +# no need to register if d4rl is install, else please register the following environments + +# register(id='door-v0', entry_point='DoorEnvV0', max_episode_steps=200, kwargs={}) +# register(id='relocate-v0', entry_point='RelocateEnvV0', max_episode_steps=200, kwargs={}) +# register(id='hammer-v0', entry_point='HammerEnvV0', max_episode_steps=200, kwargs={}) +# register(id='pen-v0', entry_point='PenEnvV0', max_episode_steps=200, kwargs={}) + +register( + id='door-shrink-finger-easy-v0', entry_point='door:DoorEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'door_shrink_finger_easy' + } +) + +register( + id='door-shrink-finger-medium-v0', entry_point='door:DoorEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'door_shrink_finger_medium' + } +) + +register( + id='door-shrink-finger-hard-v0', entry_point='door:DoorEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'door_shrink_finger_hard' + } +) + +register( + id='relocate-shrink-finger-easy-v0', entry_point='relocate:RelocateEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'relocate_shrink_finger_easy' + } +) + +register( + id='relocate-shrink-finger-medium-v0', entry_point='relocate:RelocateEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'relocate_shrink_finger_medium' + } +) + +register( + id='relocate-shrink-finger-hard-v0', entry_point='relocate:RelocateEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'relocate_shrink_finger_hard' + } +) + +register( + id='hammer-shrink-finger-easy-v0', entry_point='hammer:HammerEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'hammer_shrink_finger_easy' + } +) + +register( + id='hammer-shrink-finger-medium-v0', entry_point='hammer:HammerEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'hammer_shrink_finger_medium' + } +) + +register( + id='hammer-shrink-finger-hard-v0', entry_point='hammer:HammerEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'hammer_shrink_finger_hard' + } +) + +register( + id='pen-shrink-finger-easy-v0', entry_point='pen:PenEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'pen_shrink_finger_easy' + } +) + +register( + id='pen-shrink-finger-medium-v0', entry_point='pen:PenEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'pen_shrink_finger_medium' + } +) + +register( + id='pen-shrink-finger-hard-v0', entry_point='pen:PenEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'pen_shrink_finger_hard' + } +) + +register( + id='door-broken-joint-easy-v0', entry_point='door:DoorEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'door_broken_joint_easy' + } +) + +register( + id='door-broken-joint-medium-v0', entry_point='door:DoorEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'door_broken_joint_medium' + } +) + +register( + id='door-broken-joint-hard-v0', entry_point='door:DoorEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'door_broken_joint_hard' + } +) + +register( + id='relocate-broken-joint-easy-v0', entry_point='relocate:RelocateEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'relocate_broken_joint_easy' + } +) + +register( + id='relocate-broken-joint-medium-v0', entry_point='relocate:RelocateEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'relocate_broken_joint_medium' + } +) + +register( + id='relocate-broken-joint-hard-v0', entry_point='relocate:RelocateEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'relocate_broken_joint_hard' + } +) + +register( + id='hammer-broken-joint-easy-v0', entry_point='hammer:HammerEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'hammer_broken_joint_easy' + } +) + +register( + id='hammer-broken-joint-medium-v0', entry_point='hammer:HammerEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'hammer_broken_joint_medium' + } +) + +register( + id='hammer-broken-joint-hard-v0', entry_point='hammer:HammerEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'hammer_broken_joint_hard' + } +) + +register( + id='pen-broken-joint-easy-v0', entry_point='pen:PenEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'pen_broken_joint_easy' + } +) + +register( + id='pen-broken-joint-medium-v0', entry_point='pen:PenEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'pen_broken_joint_medium' + } +) + +register( + id='pen-broken-joint-hard-v0', entry_point='pen:PenEnvV0', max_episode_steps=200, + kwargs={ + 'xml_file': 'pen_broken_joint_hard' + } +) \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/assets/adroit.xml b/wiserl/env/odrl_envs/adroit/assets/adroit.xml new file mode 100644 index 0000000..3d00adc --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/adroit.xml @@ -0,0 +1,170 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/assets/adroit_broken_joint_easy.xml b/wiserl/env/odrl_envs/adroit/assets/adroit_broken_joint_easy.xml new file mode 100644 index 0000000..6d8c834 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/adroit_broken_joint_easy.xml @@ -0,0 +1,170 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/assets/adroit_broken_joint_hard.xml b/wiserl/env/odrl_envs/adroit/assets/adroit_broken_joint_hard.xml new file mode 100644 index 0000000..aba6805 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/adroit_broken_joint_hard.xml @@ -0,0 +1,170 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/assets/adroit_broken_joint_medium.xml b/wiserl/env/odrl_envs/adroit/assets/adroit_broken_joint_medium.xml new file mode 100644 index 0000000..d6406fc --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/adroit_broken_joint_medium.xml @@ -0,0 +1,170 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/assets/adroit_shrink_finger_easy.xml b/wiserl/env/odrl_envs/adroit/assets/adroit_shrink_finger_easy.xml new file mode 100644 index 0000000..04b236d --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/adroit_shrink_finger_easy.xml @@ -0,0 +1,171 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/assets/adroit_shrink_finger_hard.xml b/wiserl/env/odrl_envs/adroit/assets/adroit_shrink_finger_hard.xml new file mode 100644 index 0000000..847f124 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/adroit_shrink_finger_hard.xml @@ -0,0 +1,171 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/assets/adroit_shrink_finger_medium.xml b/wiserl/env/odrl_envs/adroit/assets/adroit_shrink_finger_medium.xml new file mode 100644 index 0000000..6f1349a --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/adroit_shrink_finger_medium.xml @@ -0,0 +1,171 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/assets/assets.xml b/wiserl/env/odrl_envs/adroit/assets/assets.xml new file mode 100644 index 0000000..0899c93 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/assets.xml @@ -0,0 +1,343 @@ + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/door.xml b/wiserl/env/odrl_envs/adroit/assets/door.xml new file mode 100644 index 0000000..2fedadc --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/door.xml @@ -0,0 +1,92 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/door_broken_joint_easy.xml b/wiserl/env/odrl_envs/adroit/assets/door_broken_joint_easy.xml new file mode 100644 index 0000000..ffa46e4 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/door_broken_joint_easy.xml @@ -0,0 +1,92 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/door_broken_joint_hard.xml b/wiserl/env/odrl_envs/adroit/assets/door_broken_joint_hard.xml new file mode 100644 index 0000000..5579f24 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/door_broken_joint_hard.xml @@ -0,0 +1,92 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/door_broken_joint_medium.xml b/wiserl/env/odrl_envs/adroit/assets/door_broken_joint_medium.xml new file mode 100644 index 0000000..595ae07 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/door_broken_joint_medium.xml @@ -0,0 +1,92 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/door_shrink_finger_easy.xml b/wiserl/env/odrl_envs/adroit/assets/door_shrink_finger_easy.xml new file mode 100644 index 0000000..7d24ed2 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/door_shrink_finger_easy.xml @@ -0,0 +1,92 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/door_shrink_finger_hard.xml b/wiserl/env/odrl_envs/adroit/assets/door_shrink_finger_hard.xml new file mode 100644 index 0000000..60318d7 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/door_shrink_finger_hard.xml @@ -0,0 +1,92 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/door_shrink_finger_medium.xml b/wiserl/env/odrl_envs/adroit/assets/door_shrink_finger_medium.xml new file mode 100644 index 0000000..c8ac450 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/door_shrink_finger_medium.xml @@ -0,0 +1,92 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/hammer.xml b/wiserl/env/odrl_envs/adroit/assets/hammer.xml new file mode 100644 index 0000000..560cb4b --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/hammer.xml @@ -0,0 +1,112 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/hammer_broken_joint_easy.xml b/wiserl/env/odrl_envs/adroit/assets/hammer_broken_joint_easy.xml new file mode 100644 index 0000000..bc1b196 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/hammer_broken_joint_easy.xml @@ -0,0 +1,112 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/hammer_broken_joint_hard.xml b/wiserl/env/odrl_envs/adroit/assets/hammer_broken_joint_hard.xml new file mode 100644 index 0000000..0f60577 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/hammer_broken_joint_hard.xml @@ -0,0 +1,112 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/hammer_broken_joint_medium.xml b/wiserl/env/odrl_envs/adroit/assets/hammer_broken_joint_medium.xml new file mode 100644 index 0000000..6461bfa --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/hammer_broken_joint_medium.xml @@ -0,0 +1,112 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/hammer_shrink_finger_easy.xml b/wiserl/env/odrl_envs/adroit/assets/hammer_shrink_finger_easy.xml new file mode 100644 index 0000000..6bd1bca --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/hammer_shrink_finger_easy.xml @@ -0,0 +1,112 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/hammer_shrink_finger_hard.xml b/wiserl/env/odrl_envs/adroit/assets/hammer_shrink_finger_hard.xml new file mode 100644 index 0000000..2550b06 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/hammer_shrink_finger_hard.xml @@ -0,0 +1,112 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/hammer_shrink_finger_medium.xml b/wiserl/env/odrl_envs/adroit/assets/hammer_shrink_finger_medium.xml new file mode 100644 index 0000000..84c5a91 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/hammer_shrink_finger_medium.xml @@ -0,0 +1,112 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/pen.xml b/wiserl/env/odrl_envs/adroit/assets/pen.xml new file mode 100644 index 0000000..642e5c0 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/pen.xml @@ -0,0 +1,90 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/assets/pen_broken_joint_easy.xml b/wiserl/env/odrl_envs/adroit/assets/pen_broken_joint_easy.xml new file mode 100644 index 0000000..30c17c4 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/pen_broken_joint_easy.xml @@ -0,0 +1,90 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/assets/pen_broken_joint_hard.xml b/wiserl/env/odrl_envs/adroit/assets/pen_broken_joint_hard.xml new file mode 100644 index 0000000..96cb15e --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/pen_broken_joint_hard.xml @@ -0,0 +1,90 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/assets/pen_broken_joint_medium.xml b/wiserl/env/odrl_envs/adroit/assets/pen_broken_joint_medium.xml new file mode 100644 index 0000000..72a9b86 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/pen_broken_joint_medium.xml @@ -0,0 +1,90 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/assets/pen_shrink_finger_easy.xml b/wiserl/env/odrl_envs/adroit/assets/pen_shrink_finger_easy.xml new file mode 100644 index 0000000..d709e4e --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/pen_shrink_finger_easy.xml @@ -0,0 +1,90 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/assets/pen_shrink_finger_hard.xml b/wiserl/env/odrl_envs/adroit/assets/pen_shrink_finger_hard.xml new file mode 100644 index 0000000..f0c4fcd --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/pen_shrink_finger_hard.xml @@ -0,0 +1,90 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/assets/pen_shrink_finger_medium.xml b/wiserl/env/odrl_envs/adroit/assets/pen_shrink_finger_medium.xml new file mode 100644 index 0000000..7997bf3 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/pen_shrink_finger_medium.xml @@ -0,0 +1,90 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/assets/relocate.xml b/wiserl/env/odrl_envs/adroit/assets/relocate.xml new file mode 100644 index 0000000..189caaa --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/relocate.xml @@ -0,0 +1,88 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/relocate_broken_joint_easy.xml b/wiserl/env/odrl_envs/adroit/assets/relocate_broken_joint_easy.xml new file mode 100644 index 0000000..cf33430 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/relocate_broken_joint_easy.xml @@ -0,0 +1,88 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/relocate_broken_joint_hard.xml b/wiserl/env/odrl_envs/adroit/assets/relocate_broken_joint_hard.xml new file mode 100644 index 0000000..e640baf --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/relocate_broken_joint_hard.xml @@ -0,0 +1,88 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/relocate_broken_joint_medium.xml b/wiserl/env/odrl_envs/adroit/assets/relocate_broken_joint_medium.xml new file mode 100644 index 0000000..a8d32b3 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/relocate_broken_joint_medium.xml @@ -0,0 +1,88 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/relocate_shrink_finger_easy.xml b/wiserl/env/odrl_envs/adroit/assets/relocate_shrink_finger_easy.xml new file mode 100644 index 0000000..7a2fcbd --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/relocate_shrink_finger_easy.xml @@ -0,0 +1,88 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/relocate_shrink_finger_hard.xml b/wiserl/env/odrl_envs/adroit/assets/relocate_shrink_finger_hard.xml new file mode 100644 index 0000000..03c2ba8 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/relocate_shrink_finger_hard.xml @@ -0,0 +1,88 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/assets/relocate_shrink_finger_medium.xml b/wiserl/env/odrl_envs/adroit/assets/relocate_shrink_finger_medium.xml new file mode 100644 index 0000000..cc553f8 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/assets/relocate_shrink_finger_medium.xml @@ -0,0 +1,88 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/adroit/call_adroit_env.py b/wiserl/env/odrl_envs/adroit/call_adroit_env.py new file mode 100644 index 0000000..5c10969 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/call_adroit_env.py @@ -0,0 +1,19 @@ +from typing import Dict +import gym + + +def call_adroit_env(env_config: Dict) -> gym.Env: + env_name = env_config['env_name'].lower() # eg. "pen_shrink_finger" + shift_level = env_config['shift_level'] # level(easy/medium/hard) + + if '_' in env_name: + env_name = env_name.replace('_', '-') + # decide which task it is, support the following tasks + # pen/hammer/relocate/door - shrink_finger + # - broken_joint + assert any([env_name.startswith(f'{e}') for e in ['pen', 'hammer', 'relocate', 'door']]) + assert any([env_name.endswith(f'{e}') for e in ['shrink-finger', 'broken-joint']]) + + env_name = env_name + '-' + str(shift_level) + '-v0' + + return gym.make(env_name) \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/F1.stl b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/F1.stl new file mode 100644 index 0000000..515d3c9 Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/F1.stl differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/F2.stl b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/F2.stl new file mode 100644 index 0000000..7bc5e20 Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/F2.stl differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/F3.stl b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/F3.stl new file mode 100644 index 0000000..223f06f Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/F3.stl differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/TH1_z.stl b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/TH1_z.stl new file mode 100644 index 0000000..400ee2d Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/TH1_z.stl differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/TH2_z.stl b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/TH2_z.stl new file mode 100644 index 0000000..5ace838 Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/TH2_z.stl differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/TH3_z.stl b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/TH3_z.stl new file mode 100644 index 0000000..23485ab Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/TH3_z.stl differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/forearm_simple.stl b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/forearm_simple.stl new file mode 100644 index 0000000..888d2d3 Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/forearm_simple.stl differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/knuckle.stl b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/knuckle.stl new file mode 100644 index 0000000..4faedd7 Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/knuckle.stl differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/lfmetacarpal.stl b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/lfmetacarpal.stl new file mode 100644 index 0000000..535cf4d Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/lfmetacarpal.stl differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/palm.stl b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/palm.stl new file mode 100644 index 0000000..65e47eb Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/palm.stl differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/wrist.stl b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/wrist.stl new file mode 100644 index 0000000..420d5f9 Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/meshes/wrist.stl differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/darkwood.png b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/darkwood.png new file mode 100644 index 0000000..d5dcc5c Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/darkwood.png differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/dice.png b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/dice.png new file mode 100644 index 0000000..798a8e0 Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/dice.png differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/foil.png b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/foil.png new file mode 100644 index 0000000..654cfe1 Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/foil.png differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/marble.png b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/marble.png new file mode 100644 index 0000000..c50e8b9 Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/marble.png differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/silverRaw.png b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/silverRaw.png new file mode 100644 index 0000000..13690e5 Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/silverRaw.png differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/skin.png b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/skin.png new file mode 100644 index 0000000..54e528d Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/skin.png differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/square.png b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/square.png new file mode 100644 index 0000000..dbfd695 Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/square.png differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/wood.png b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/wood.png new file mode 100644 index 0000000..c323cb9 Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/wood.png differ diff --git a/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/woodb.png b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/woodb.png new file mode 100644 index 0000000..47f94a8 Binary files /dev/null and b/wiserl/env/odrl_envs/adroit/dependencies/Adroit/resources/textures/woodb.png differ diff --git a/wiserl/env/odrl_envs/adroit/door.py b/wiserl/env/odrl_envs/adroit/door.py new file mode 100644 index 0000000..f075ec4 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/door.py @@ -0,0 +1,126 @@ +import numpy as np +from gym import utils +from mujoco_py import MjViewer +import os +import mujoco_env + +ADD_BONUS_REWARDS = True + +class DoorEnvV0(mujoco_env.MujocoEnv, utils.EzPickle): + def __init__(self, xml_file: str = "door"): + self.door_hinge_did = 0 + self.door_bid = 0 + self.grasp_sid = 0 + self.handle_sid = 0 + curr_dir = os.path.dirname(os.path.abspath(__file__)) + mujoco_env.MujocoEnv.__init__(self, curr_dir + f'/assets/{xml_file}.xml', 5) + + # change actuator sensitivity + self.sim.model.actuator_gainprm[self.sim.model.actuator_name2id('A_WRJ1'):self.sim.model.actuator_name2id('A_WRJ0')+1,:3] = np.array([10, 0, 0]) + self.sim.model.actuator_gainprm[self.sim.model.actuator_name2id('A_FFJ3'):self.sim.model.actuator_name2id('A_THJ0')+1,:3] = np.array([1, 0, 0]) + self.sim.model.actuator_biasprm[self.sim.model.actuator_name2id('A_WRJ1'):self.sim.model.actuator_name2id('A_WRJ0')+1,:3] = np.array([0, -10, 0]) + self.sim.model.actuator_biasprm[self.sim.model.actuator_name2id('A_FFJ3'):self.sim.model.actuator_name2id('A_THJ0')+1,:3] = np.array([0, -1, 0]) + + utils.EzPickle.__init__(self) + ob = self.reset_model() + self.act_mid = np.mean(self.model.actuator_ctrlrange, axis=1) + self.act_rng = 0.5*(self.model.actuator_ctrlrange[:,1]-self.model.actuator_ctrlrange[:,0]) + self.action_space.high = np.ones_like(self.model.actuator_ctrlrange[:,1]) + self.action_space.low = -1.0 * np.ones_like(self.model.actuator_ctrlrange[:,0]) + self.door_hinge_did = self.model.jnt_dofadr[self.model.joint_name2id('door_hinge')] + self.grasp_sid = self.model.site_name2id('S_grasp') + self.handle_sid = self.model.site_name2id('S_handle') + self.door_bid = self.model.body_name2id('frame') + + def step(self, a): + a = np.clip(a, -1.0, 1.0) + try: + a = self.act_mid + a*self.act_rng # mean center and scale + except: + a = a # only for the initialization phase + self.do_simulation(a, self.frame_skip) + ob = self.get_obs() + handle_pos = self.data.site_xpos[self.handle_sid].ravel() + palm_pos = self.data.site_xpos[self.grasp_sid].ravel() + door_pos = self.data.qpos[self.door_hinge_did] + + # get to handle + reward = -0.1*np.linalg.norm(palm_pos-handle_pos) + # open door + reward += -0.1*(door_pos - 1.57)*(door_pos - 1.57) + # velocity cost + reward += -1e-5*np.sum(self.data.qvel**2) + + if ADD_BONUS_REWARDS: + # Bonus + if door_pos > 0.2: + reward += 2 + if door_pos > 1.0: + reward += 8 + if door_pos > 1.35: + reward += 10 + + goal_achieved = True if door_pos >= 1.35 else False + + return ob, reward, False, dict(goal_achieved=goal_achieved) + + def get_obs(self): + # qpos for hand + # xpos for obj + # xpos for target + qp = self.data.qpos.ravel() + handle_pos = self.data.site_xpos[self.handle_sid].ravel() + palm_pos = self.data.site_xpos[self.grasp_sid].ravel() + door_pos = np.array([self.data.qpos[self.door_hinge_did]]) + if door_pos > 1.0: + door_open = 1.0 + else: + door_open = -1.0 + latch_pos = qp[-1] + return np.concatenate([qp[1:-2], [latch_pos], door_pos, palm_pos, handle_pos, palm_pos-handle_pos, [door_open]]) + + def reset_model(self): + qp = self.init_qpos.copy() + qv = self.init_qvel.copy() + self.set_state(qp, qv) + + self.model.body_pos[self.door_bid,0] = self.np_random.uniform(low=-0.3, high=-0.2) + self.model.body_pos[self.door_bid,1] = self.np_random.uniform(low=0.25, high=0.35) + self.model.body_pos[self.door_bid,2] = self.np_random.uniform(low=0.252, high=0.35) + self.sim.forward() + return self.get_obs() + + def get_env_state(self): + """ + Get state of hand as well as objects and targets in the scene + """ + qp = self.data.qpos.ravel().copy() + qv = self.data.qvel.ravel().copy() + door_body_pos = self.model.body_pos[self.door_bid].ravel().copy() + return dict(qpos=qp, qvel=qv, door_body_pos=door_body_pos) + + def set_env_state(self, state_dict): + """ + Set the state which includes hand as well as objects and targets in the scene + """ + qp = state_dict['qpos'] + qv = state_dict['qvel'] + self.set_state(qp, qv) + self.model.body_pos[self.door_bid] = state_dict['door_body_pos'] + self.sim.forward() + + def mj_viewer_setup(self): + self.viewer = MjViewer(self.sim) + self.viewer.cam.azimuth = 90 + self.sim.forward() + self.viewer.cam.distance = 1.5 + + def evaluate_success(self, paths): + num_success = 0 + num_paths = len(paths) + # success if door open for 25 steps + for path in paths: + if np.sum(path['env_infos']['goal_achieved']) > 25: + num_success += 1 + success_percentage = num_success*100.0/num_paths + return success_percentage diff --git a/wiserl/env/odrl_envs/adroit/hammer.py b/wiserl/env/odrl_envs/adroit/hammer.py new file mode 100644 index 0000000..b299a08 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/hammer.py @@ -0,0 +1,132 @@ +import numpy as np +from gym import utils +from mujoco_py import MjViewer +import os + +from quatmath import quat2euler + +import mujoco_env + +ADD_BONUS_REWARDS = True + +class HammerEnvV0(mujoco_env.MujocoEnv, utils.EzPickle): + def __init__(self, xml_file: str = 'hammer'): + self.target_obj_sid = -1 + self.S_grasp_sid = -1 + self.obj_bid = -1 + self.tool_sid = -1 + self.goal_sid = -1 + curr_dir = os.path.dirname(os.path.abspath(__file__)) + mujoco_env.MujocoEnv.__init__(self, curr_dir + f'/assets/{xml_file}.xml', 5) + utils.EzPickle.__init__(self) + + # change actuator sensitivity + self.sim.model.actuator_gainprm[self.sim.model.actuator_name2id('A_WRJ1'):self.sim.model.actuator_name2id('A_WRJ0')+1,:3] = np.array([10, 0, 0]) + self.sim.model.actuator_gainprm[self.sim.model.actuator_name2id('A_FFJ3'):self.sim.model.actuator_name2id('A_THJ0')+1,:3] = np.array([1, 0, 0]) + self.sim.model.actuator_biasprm[self.sim.model.actuator_name2id('A_WRJ1'):self.sim.model.actuator_name2id('A_WRJ0')+1,:3] = np.array([0, -10, 0]) + self.sim.model.actuator_biasprm[self.sim.model.actuator_name2id('A_FFJ3'):self.sim.model.actuator_name2id('A_THJ0')+1,:3] = np.array([0, -1, 0]) + + self.target_obj_sid = self.sim.model.site_name2id('S_target') + self.S_grasp_sid = self.sim.model.site_name2id('S_grasp') + self.obj_bid = self.sim.model.body_name2id('Object') + self.tool_sid = self.sim.model.site_name2id('tool') + self.goal_sid = self.sim.model.site_name2id('nail_goal') + self.act_mid = np.mean(self.model.actuator_ctrlrange, axis=1) + self.act_rng = 0.5 * (self.model.actuator_ctrlrange[:, 1] - self.model.actuator_ctrlrange[:, 0]) + self.action_space.high = np.ones_like(self.model.actuator_ctrlrange[:,1]) + self.action_space.low = -1.0 * np.ones_like(self.model.actuator_ctrlrange[:,0]) + + def step(self, a): + a = np.clip(a, -1.0, 1.0) + try: + a = self.act_mid + a * self.act_rng # mean center and scale + except: + a = a # only for the initialization phase + self.do_simulation(a, self.frame_skip) + ob = self.get_obs() + obj_pos = self.data.body_xpos[self.obj_bid].ravel() + palm_pos = self.data.site_xpos[self.S_grasp_sid].ravel() + tool_pos = self.data.site_xpos[self.tool_sid].ravel() + target_pos = self.data.site_xpos[self.target_obj_sid].ravel() + goal_pos = self.data.site_xpos[self.goal_sid].ravel() + + # get to hammer + reward = - 0.1 * np.linalg.norm(palm_pos - obj_pos) + # take hammer head to nail + reward -= np.linalg.norm((tool_pos - target_pos)) + # make nail go inside + reward -= 10 * np.linalg.norm(target_pos - goal_pos) + # velocity penalty + reward -= 1e-2 * np.linalg.norm(self.data.qvel.ravel()) + + if ADD_BONUS_REWARDS: + # bonus for lifting up the hammer + if obj_pos[2] > 0.04 and tool_pos[2] > 0.04: + reward += 2 + + # bonus for hammering the nail + if (np.linalg.norm(target_pos - goal_pos) < 0.020): + reward += 25 + if (np.linalg.norm(target_pos - goal_pos) < 0.010): + reward += 75 + + goal_achieved = True if np.linalg.norm(target_pos - goal_pos) < 0.010 else False + + return ob, reward, False, dict(goal_achieved=goal_achieved) + + def get_obs(self): + # qpos for hand + # xpos for obj + # xpos for target + qp = self.data.qpos.ravel() + qv = np.clip(self.data.qvel.ravel(), -1.0, 1.0) + obj_pos = self.data.body_xpos[self.obj_bid].ravel() + obj_rot = quat2euler(self.data.body_xquat[self.obj_bid].ravel()).ravel() + palm_pos = self.data.site_xpos[self.S_grasp_sid].ravel() + target_pos = self.data.site_xpos[self.target_obj_sid].ravel() + nail_impact = 0.0 + return np.concatenate([qp[:-6], qv[-6:], palm_pos, obj_pos, obj_rot, target_pos, np.array([nail_impact])]) + + def reset_model(self): + self.sim.reset() + target_bid = self.model.body_name2id('nail_board') + self.model.body_pos[target_bid,2] = self.np_random.uniform(low=0.1, high=0.25) + self.sim.forward() + return self.get_obs() + + def get_env_state(self): + """ + Get state of hand as well as objects and targets in the scene + """ + qpos = self.data.qpos.ravel().copy() + qvel = self.data.qvel.ravel().copy() + board_pos = self.model.body_pos[self.model.body_name2id('nail_board')].copy() + target_pos = self.data.site_xpos[self.target_obj_sid].ravel().copy() + return dict(qpos=qpos, qvel=qvel, board_pos=board_pos, target_pos=target_pos) + + def set_env_state(self, state_dict): + """ + Set the state which includes hand as well as objects and targets in the scene + """ + qp = state_dict['qpos'] + qv = state_dict['qvel'] + board_pos = state_dict['board_pos'] + self.set_state(qp, qv) + self.model.body_pos[self.model.body_name2id('nail_board')] = board_pos + self.sim.forward() + + def mj_viewer_setup(self): + self.viewer = MjViewer(self.sim) + self.viewer.cam.azimuth = 45 + self.viewer.cam.distance = 2.0 + self.sim.forward() + + def evaluate_success(self, paths): + num_success = 0 + num_paths = len(paths) + # success if nail insude board for 25 steps + for path in paths: + if np.sum(path['env_infos']['goal_achieved']) > 25: + num_success += 1 + success_percentage = num_success*100.0/num_paths + return success_percentage diff --git a/wiserl/env/odrl_envs/adroit/mujoco_env.py b/wiserl/env/odrl_envs/adroit/mujoco_env.py new file mode 100644 index 0000000..8b6967c --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/mujoco_env.py @@ -0,0 +1,142 @@ +import os + +from gym import error, spaces +from gym.utils import seeding +import numpy as np +from os import path +import gym +import six +import time as timer + +try: + import mujoco_py + from mujoco_py import load_model_from_path, load_model_from_xml, MjSim, MjViewer +except ImportError as e: + raise error.DependencyNotInstalled("{}. (HINT: you need to install mujoco_py, and also perform the setup instructions here: https://github.com/openai/mujoco-py/.)".format(e)) + +def get_sim(model_path, model_xml): + if model_xml is None: + if model_path.startswith("/"): + fullpath = model_path + else: + fullpath = os.path.join(os.path.dirname(__file__), "assets", model_path) + if not path.exists(fullpath): + raise IOError("File %s does not exist" % fullpath) + model = load_model_from_path(fullpath) + else: + model = load_model_from_xml(model_xml) + return MjSim(model) + +class MujocoEnv(gym.Env): + """Superclass for all MuJoCo environments. + """ + + def __init__(self, model_path, frame_skip=1, model_xml=None, sim=None): + + if sim is None: + self.sim = get_sim(model_path=model_path, model_xml=model_xml) + else: + self.sim = sim + self.data = self.sim.data + self.model = self.sim.model + + self.frame_skip = frame_skip + self.metadata = { + 'render.modes': ['human', 'rgb_array'], + 'video.frames_per_second': int(np.round(1.0 / self.dt)) + } + self.mujoco_render_frames = False + + self.init_qpos = self.data.qpos.ravel().copy() + self.init_qvel = self.data.qvel.ravel().copy() + try: + observation, _reward, done, _info = self.step(np.zeros(self.model.nu)) + except NotImplementedError: + observation, _reward, done, _info = self._step(np.zeros(self.model.nu)) + assert not done + self.obs_dim = np.sum([o.size for o in observation]) if type(observation) is tuple else observation.size + + bounds = self.model.actuator_ctrlrange.copy() + low = bounds[:, 0] + high = bounds[:, 1] + self.action_space = spaces.Box(low, high, dtype=np.float64) + + high = np.inf*np.ones(self.obs_dim) + low = -high + self.observation_space = spaces.Box(low, high, dtype=np.float64) + + self.seed() + + def seed(self, seed=None): + self.np_random, seed = seeding.np_random(seed) + return [seed] + + # methods to override: + # ---------------------------- + + def reset_model(self): + """ + Reset the robot degrees of freedom (qpos and qvel). + Implement this in each subclass. + """ + raise NotImplementedError + + def mj_viewer_setup(self): + """ + Due to specifics of new mujoco rendering, the standard viewer cannot be used + with this set-up. Instead we use this mujoco specific function. + """ + pass + + def viewer_setup(self): + """ + Does not work. Use mj_viewer_setup() instead + """ + pass + + def evaluate_success(self, paths, logger=None): + """ + Log various success metrics calculated based on input paths into the logger + """ + pass + + # ----------------------------- + + def reset(self): + self.sim.reset() + self.sim.forward() + ob = self.reset_model() + return ob + + def set_state(self, qpos, qvel): + assert qpos.shape == (self.model.nq,) and qvel.shape == (self.model.nv,) + old_state = self.sim.get_state() + new_state = mujoco_py.MjSimState(old_state.time, qpos, qvel, + old_state.act, old_state.udd_state) + self.sim.set_state(new_state) + self.sim.forward() + + @property + def dt(self): + return self.model.opt.timestep * self.frame_skip + + def do_simulation(self, ctrl, n_frames): + for i in range(self.model.nu): + self.sim.data.ctrl[i] = ctrl[i] + for _ in range(n_frames): + self.sim.step() + if self.mujoco_render_frames is True: + self.mj_render() + + def mj_render(self): + try: + self.viewer.render() + except: + self.mj_viewer_setup() + self.viewer._run_speed = 0.5 + self.viewer._run_speed /= self.frame_skip + self.viewer.render() + + def render(self, *args, **kwargs): + self.mj_render() + diff --git a/wiserl/env/odrl_envs/adroit/pen.py b/wiserl/env/odrl_envs/adroit/pen.py new file mode 100644 index 0000000..31989f6 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/pen.py @@ -0,0 +1,147 @@ +import numpy as np +from gym import utils +from gym import spaces +from mujoco_py import MjViewer +import os + +import mujoco_env +from quatmath import euler2quat + +ADD_BONUS_REWARDS = True + +class PenEnvV0(mujoco_env.MujocoEnv, utils.EzPickle): + def __init__(self, xml_file: str = 'pen'): + self.target_obj_bid = 0 + self.S_grasp_sid = 0 + self.eps_ball_sid = 0 + self.obj_bid = 0 + self.obj_t_sid = 0 + self.obj_b_sid = 0 + self.tar_t_sid = 0 + self.tar_b_sid = 0 + self.pen_length = 1.0 + self.tar_length = 1.0 + + curr_dir = os.path.dirname(os.path.abspath(__file__)) + mujoco_env.MujocoEnv.__init__(self, curr_dir+ f'/assets/{xml_file}.xml', 5) + + # Override action_space to -1, 1 + self.action_space = spaces.Box(low=-1.0, high=1.0, dtype=np.float32, shape=self.action_space.shape) + + # change actuator sensitivity + self.sim.model.actuator_gainprm[self.sim.model.actuator_name2id('A_WRJ1'):self.sim.model.actuator_name2id('A_WRJ0')+1,:3] = np.array([10, 0, 0]) + self.sim.model.actuator_gainprm[self.sim.model.actuator_name2id('A_FFJ3'):self.sim.model.actuator_name2id('A_THJ0')+1,:3] = np.array([1, 0, 0]) + self.sim.model.actuator_biasprm[self.sim.model.actuator_name2id('A_WRJ1'):self.sim.model.actuator_name2id('A_WRJ0')+1,:3] = np.array([0, -10, 0]) + self.sim.model.actuator_biasprm[self.sim.model.actuator_name2id('A_FFJ3'):self.sim.model.actuator_name2id('A_THJ0')+1,:3] = np.array([0, -1, 0]) + + utils.EzPickle.__init__(self) + self.target_obj_bid = self.sim.model.body_name2id("target") + self.S_grasp_sid = self.sim.model.site_name2id('S_grasp') + self.obj_bid = self.sim.model.body_name2id('Object') + self.eps_ball_sid = self.sim.model.site_name2id('eps_ball') + self.obj_t_sid = self.sim.model.site_name2id('object_top') + self.obj_b_sid = self.sim.model.site_name2id('object_bottom') + self.tar_t_sid = self.sim.model.site_name2id('target_top') + self.tar_b_sid = self.sim.model.site_name2id('target_bottom') + + self.pen_length = np.linalg.norm(self.data.site_xpos[self.obj_t_sid] - self.data.site_xpos[self.obj_b_sid]) + self.tar_length = np.linalg.norm(self.data.site_xpos[self.tar_t_sid] - self.data.site_xpos[self.tar_b_sid]) + + self.act_mid = np.mean(self.model.actuator_ctrlrange, axis=1) + self.act_rng = 0.5*(self.model.actuator_ctrlrange[:,1]-self.model.actuator_ctrlrange[:,0]) + + def step(self, a): + a = np.clip(a, -1.0, 1.0) + try: + starting_up = False + a = self.act_mid + a*self.act_rng # mean center and scale + except: + starting_up = True + a = a # only for the initialization phase + self.do_simulation(a, self.frame_skip) + + obj_pos = self.data.body_xpos[self.obj_bid].ravel() + desired_loc = self.data.site_xpos[self.eps_ball_sid].ravel() + obj_orien = (self.data.site_xpos[self.obj_t_sid] - self.data.site_xpos[self.obj_b_sid])/self.pen_length + desired_orien = (self.data.site_xpos[self.tar_t_sid] - self.data.site_xpos[self.tar_b_sid])/self.tar_length + + # pos cost + dist = np.linalg.norm(obj_pos-desired_loc) + reward = -dist + # orien cost + orien_similarity = np.dot(obj_orien, desired_orien) + reward += orien_similarity + + if ADD_BONUS_REWARDS: + # bonus for being close to desired orientation + if dist < 0.075 and orien_similarity > 0.9: + reward += 10 + if dist < 0.075 and orien_similarity > 0.95: + reward += 50 + + # penalty for dropping the pen + done = False + if obj_pos[2] < 0.075: + reward -= 5 + done = True if not starting_up else False + + goal_achieved = True if (dist < 0.075 and orien_similarity > 0.95) else False + + return self.get_obs(), reward, done, dict(goal_achieved=goal_achieved) + + def get_obs(self): + qp = self.data.qpos.ravel() + obj_vel = self.data.qvel[-6:].ravel() + obj_pos = self.data.body_xpos[self.obj_bid].ravel() + desired_pos = self.data.site_xpos[self.eps_ball_sid].ravel() + obj_orien = (self.data.site_xpos[self.obj_t_sid] - self.data.site_xpos[self.obj_b_sid])/self.pen_length + desired_orien = (self.data.site_xpos[self.tar_t_sid] - self.data.site_xpos[self.tar_b_sid])/self.tar_length + return np.concatenate([qp[:-6], obj_pos, obj_vel, obj_orien, desired_orien, + obj_pos-desired_pos, obj_orien-desired_orien]) + + def reset_model(self): + qp = self.init_qpos.copy() + qv = self.init_qvel.copy() + self.set_state(qp, qv) + desired_orien = np.zeros(3) + desired_orien[0] = self.np_random.uniform(low=-1, high=1) + desired_orien[1] = self.np_random.uniform(low=-1, high=1) + self.model.body_quat[self.target_obj_bid] = euler2quat(desired_orien) + self.sim.forward() + return self.get_obs() + + def get_env_state(self): + """ + Get state of hand as well as objects and targets in the scene + """ + qp = self.data.qpos.ravel().copy() + qv = self.data.qvel.ravel().copy() + desired_orien = self.model.body_quat[self.target_obj_bid].ravel().copy() + return dict(qpos=qp, qvel=qv, desired_orien=desired_orien) + + def set_env_state(self, state_dict): + """ + Set the state which includes hand as well as objects and targets in the scene + """ + qp = state_dict['qpos'] + qv = state_dict['qvel'] + desired_orien = state_dict['desired_orien'] + self.set_state(qp, qv) + self.model.body_quat[self.target_obj_bid] = desired_orien + self.sim.forward() + + def mj_viewer_setup(self): + self.viewer = MjViewer(self.sim) + self.viewer.cam.azimuth = -45 + self.sim.forward() + self.viewer.cam.distance = 1.0 + + def evaluate_success(self, paths): + num_success = 0 + num_paths = len(paths) + # success if pen within 15 degrees of target for 20 steps + for path in paths: + if np.sum(path['env_infos']['goal_achieved']) > 20: + num_success += 1 + success_percentage = num_success*100.0/num_paths + return success_percentage \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/quatmath.py b/wiserl/env/odrl_envs/adroit/quatmath.py new file mode 100644 index 0000000..7fef129 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/quatmath.py @@ -0,0 +1,164 @@ +import numpy as np +# For testing whether a number is close to zero +_FLOAT_EPS = np.finfo(np.float64).eps +_EPS4 = _FLOAT_EPS * 4.0 + + +def mulQuat(qa, qb): + res = np.zeros(4) + res[0] = qa[0]*qb[0] - qa[1]*qb[1] - qa[2]*qb[2] - qa[3]*qb[3] + res[1] = qa[0]*qb[1] + qa[1]*qb[0] + qa[2]*qb[3] - qa[3]*qb[2] + res[2] = qa[0]*qb[2] - qa[1]*qb[3] + qa[2]*qb[0] + qa[3]*qb[1] + res[3] = qa[0]*qb[3] + qa[1]*qb[2] - qa[2]*qb[1] + qa[3]*qb[0] + return res + +def negQuat(quat): + return np.array([quat[0], -quat[1], -quat[2], -quat[3]]) + +def quat2Vel(quat, dt=1): + axis = quat[1:].copy() + sin_a_2 = np.sqrt(np.sum(axis**2)) + axis = axis/(sin_a_2+1e-8) + speed = 2*np.arctan2(sin_a_2, quat[0])/dt + return speed, axis + +def quatDiff2Vel(quat1, quat2, dt): + neg = negQuat(quat1) + diff = mulQuat(quat2, neg) + return quat2Vel(diff, dt) + + +def axis_angle2quat(axis, angle): + c = np.cos(angle/2) + s = np.sin(angle/2) + return np.array([c, s*axis[0], s*axis[1], s*axis[2]]) + +def euler2mat(euler): + """ Convert Euler Angles to Rotation Matrix. See rotation.py for notes """ + euler = np.asarray(euler, dtype=np.float64) + assert euler.shape[-1] == 3, "Invalid shaped euler {}".format(euler) + + ai, aj, ak = -euler[..., 2], -euler[..., 1], -euler[..., 0] + si, sj, sk = np.sin(ai), np.sin(aj), np.sin(ak) + ci, cj, ck = np.cos(ai), np.cos(aj), np.cos(ak) + cc, cs = ci * ck, ci * sk + sc, ss = si * ck, si * sk + + mat = np.empty(euler.shape[:-1] + (3, 3), dtype=np.float64) + mat[..., 2, 2] = cj * ck + mat[..., 2, 1] = sj * sc - cs + mat[..., 2, 0] = sj * cc + ss + mat[..., 1, 2] = cj * sk + mat[..., 1, 1] = sj * ss + cc + mat[..., 1, 0] = sj * cs - sc + mat[..., 0, 2] = -sj + mat[..., 0, 1] = cj * si + mat[..., 0, 0] = cj * ci + return mat + + +def euler2quat(euler): + """ Convert Euler Angles to Quaternions. See rotation.py for notes """ + euler = np.asarray(euler, dtype=np.float64) + assert euler.shape[-1] == 3, "Invalid shape euler {}".format(euler) + + ai, aj, ak = euler[..., 2] / 2, -euler[..., 1] / 2, euler[..., 0] / 2 + si, sj, sk = np.sin(ai), np.sin(aj), np.sin(ak) + ci, cj, ck = np.cos(ai), np.cos(aj), np.cos(ak) + cc, cs = ci * ck, ci * sk + sc, ss = si * ck, si * sk + + quat = np.empty(euler.shape[:-1] + (4,), dtype=np.float64) + quat[..., 0] = cj * cc + sj * ss + quat[..., 3] = cj * sc - sj * cs + quat[..., 2] = -(cj * ss + sj * cc) + quat[..., 1] = cj * cs - sj * sc + return quat + + +def mat2euler(mat): + """ Convert Rotation Matrix to Euler Angles. See rotation.py for notes """ + mat = np.asarray(mat, dtype=np.float64) + assert mat.shape[-2:] == (3, 3), "Invalid shape matrix {}".format(mat) + + cy = np.sqrt(mat[..., 2, 2] * mat[..., 2, 2] + mat[..., 1, 2] * mat[..., 1, 2]) + condition = cy > _EPS4 + euler = np.empty(mat.shape[:-1], dtype=np.float64) + euler[..., 2] = np.where(condition, + -np.arctan2(mat[..., 0, 1], mat[..., 0, 0]), + -np.arctan2(-mat[..., 1, 0], mat[..., 1, 1])) + euler[..., 1] = np.where(condition, + -np.arctan2(-mat[..., 0, 2], cy), + -np.arctan2(-mat[..., 0, 2], cy)) + euler[..., 0] = np.where(condition, + -np.arctan2(mat[..., 1, 2], mat[..., 2, 2]), + 0.0) + return euler + + +def mat2quat(mat): + """ Convert Rotation Matrix to Quaternion. See rotation.py for notes """ + mat = np.asarray(mat, dtype=np.float64) + assert mat.shape[-2:] == (3, 3), "Invalid shape matrix {}".format(mat) + + Qxx, Qyx, Qzx = mat[..., 0, 0], mat[..., 0, 1], mat[..., 0, 2] + Qxy, Qyy, Qzy = mat[..., 1, 0], mat[..., 1, 1], mat[..., 1, 2] + Qxz, Qyz, Qzz = mat[..., 2, 0], mat[..., 2, 1], mat[..., 2, 2] + # Fill only lower half of symmetric matrix + K = np.zeros(mat.shape[:-2] + (4, 4), dtype=np.float64) + K[..., 0, 0] = Qxx - Qyy - Qzz + K[..., 1, 0] = Qyx + Qxy + K[..., 1, 1] = Qyy - Qxx - Qzz + K[..., 2, 0] = Qzx + Qxz + K[..., 2, 1] = Qzy + Qyz + K[..., 2, 2] = Qzz - Qxx - Qyy + K[..., 3, 0] = Qyz - Qzy + K[..., 3, 1] = Qzx - Qxz + K[..., 3, 2] = Qxy - Qyx + K[..., 3, 3] = Qxx + Qyy + Qzz + K /= 3.0 + # TODO: vectorize this -- probably could be made faster + q = np.empty(K.shape[:-2] + (4,)) + it = np.nditer(q[..., 0], flags=['multi_index']) + while not it.finished: + # Use Hermitian eigenvectors, values for speed + vals, vecs = np.linalg.eigh(K[it.multi_index]) + # Select largest eigenvector, reorder to w,x,y,z quaternion + q[it.multi_index] = vecs[[3, 0, 1, 2], np.argmax(vals)] + # Prefer quaternion with positive w + # (q * -1 corresponds to same rotation as q) + if q[it.multi_index][0] < 0: + q[it.multi_index] *= -1 + it.iternext() + return q + + +def quat2euler(quat): + """ Convert Quaternion to Euler Angles. See rotation.py for notes """ + return mat2euler(quat2mat(quat)) + + +def quat2mat(quat): + """ Convert Quaternion to Euler Angles. See rotation.py for notes """ + quat = np.asarray(quat, dtype=np.float64) + assert quat.shape[-1] == 4, "Invalid shape quat {}".format(quat) + + w, x, y, z = quat[..., 0], quat[..., 1], quat[..., 2], quat[..., 3] + Nq = np.sum(quat * quat, axis=-1) + s = 2.0 / Nq + X, Y, Z = x * s, y * s, z * s + wX, wY, wZ = w * X, w * Y, w * Z + xX, xY, xZ = x * X, x * Y, x * Z + yY, yZ, zZ = y * Y, y * Z, z * Z + + mat = np.empty(quat.shape[:-1] + (3, 3), dtype=np.float64) + mat[..., 0, 0] = 1.0 - (yY + zZ) + mat[..., 0, 1] = xY - wZ + mat[..., 0, 2] = xZ + wY + mat[..., 1, 0] = xY + wZ + mat[..., 1, 1] = 1.0 - (xX + zZ) + mat[..., 1, 2] = yZ - wX + mat[..., 2, 0] = xZ - wY + mat[..., 2, 1] = yZ + wX + mat[..., 2, 2] = 1.0 - (xX + yY) + return np.where((Nq > _FLOAT_EPS)[..., np.newaxis, np.newaxis], mat, np.eye(3)) \ No newline at end of file diff --git a/wiserl/env/odrl_envs/adroit/relocate.py b/wiserl/env/odrl_envs/adroit/relocate.py new file mode 100644 index 0000000..641a753 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/relocate.py @@ -0,0 +1,123 @@ +import numpy as np +from gym import utils +from mujoco_py import MjViewer +import os + +import mujoco_env + +ADD_BONUS_REWARDS = True + +class RelocateEnvV0(mujoco_env.MujocoEnv, utils.EzPickle): + def __init__(self, xml_file: str = 'relocate'): + self.target_obj_sid = 0 + self.S_grasp_sid = 0 + self.obj_bid = 0 + curr_dir = os.path.dirname(os.path.abspath(__file__)) + mujoco_env.MujocoEnv.__init__(self, curr_dir + f'/assets/{xml_file}.xml', 5) + + # change actuator sensitivity + self.sim.model.actuator_gainprm[self.sim.model.actuator_name2id('A_WRJ1'):self.sim.model.actuator_name2id('A_WRJ0')+1,:3] = np.array([10, 0, 0]) + self.sim.model.actuator_gainprm[self.sim.model.actuator_name2id('A_FFJ3'):self.sim.model.actuator_name2id('A_THJ0')+1,:3] = np.array([1, 0, 0]) + self.sim.model.actuator_biasprm[self.sim.model.actuator_name2id('A_WRJ1'):self.sim.model.actuator_name2id('A_WRJ0')+1,:3] = np.array([0, -10, 0]) + self.sim.model.actuator_biasprm[self.sim.model.actuator_name2id('A_FFJ3'):self.sim.model.actuator_name2id('A_THJ0')+1,:3] = np.array([0, -1, 0]) + + self.target_obj_sid = self.sim.model.site_name2id("target") + self.S_grasp_sid = self.sim.model.site_name2id('S_grasp') + self.obj_bid = self.sim.model.body_name2id('Object') + utils.EzPickle.__init__(self) + self.act_mid = np.mean(self.model.actuator_ctrlrange, axis=1) + self.act_rng = 0.5*(self.model.actuator_ctrlrange[:,1]-self.model.actuator_ctrlrange[:,0]) + self.action_space.high = np.ones_like(self.model.actuator_ctrlrange[:,1]) + self.action_space.low = -1.0 * np.ones_like(self.model.actuator_ctrlrange[:,0]) + + def step(self, a): + a = np.clip(a, -1.0, 1.0) + try: + a = self.act_mid + a*self.act_rng # mean center and scale + except: + a = a # only for the initialization phase + self.do_simulation(a, self.frame_skip) + ob = self.get_obs() + obj_pos = self.data.body_xpos[self.obj_bid].ravel() + palm_pos = self.data.site_xpos[self.S_grasp_sid].ravel() + target_pos = self.data.site_xpos[self.target_obj_sid].ravel() + + reward = -0.1*np.linalg.norm(palm_pos-obj_pos) # take hand to object + if obj_pos[2] > 0.04: # if object off the table + reward += 1.0 # bonus for lifting the object + reward += -0.5*np.linalg.norm(palm_pos-target_pos) # make hand go to target + reward += -0.5*np.linalg.norm(obj_pos-target_pos) # make object go to target + + if ADD_BONUS_REWARDS: + if np.linalg.norm(obj_pos-target_pos) < 0.1: + reward += 10.0 # bonus for object close to target + if np.linalg.norm(obj_pos-target_pos) < 0.05: + reward += 20.0 # bonus for object "very" close to target + + goal_achieved = True if np.linalg.norm(obj_pos-target_pos) < 0.1 else False + + return ob, reward, False, dict(goal_achieved=goal_achieved) + + def get_obs(self): + # qpos for hand + # xpos for obj + # xpos for target + qp = self.data.qpos.ravel() + obj_pos = self.data.body_xpos[self.obj_bid].ravel() + palm_pos = self.data.site_xpos[self.S_grasp_sid].ravel() + target_pos = self.data.site_xpos[self.target_obj_sid].ravel() + return np.concatenate([qp[:-6], palm_pos-obj_pos, palm_pos-target_pos, obj_pos-target_pos]) + + def reset_model(self): + qp = self.init_qpos.copy() + qv = self.init_qvel.copy() + self.set_state(qp, qv) + self.model.body_pos[self.obj_bid,0] = self.np_random.uniform(low=-0.15, high=0.15) + self.model.body_pos[self.obj_bid,1] = self.np_random.uniform(low=-0.15, high=0.3) + self.model.site_pos[self.target_obj_sid, 0] = self.np_random.uniform(low=-0.2, high=0.2) + self.model.site_pos[self.target_obj_sid,1] = self.np_random.uniform(low=-0.2, high=0.2) + self.model.site_pos[self.target_obj_sid,2] = self.np_random.uniform(low=0.15, high=0.35) + self.sim.forward() + return self.get_obs() + + def get_env_state(self): + """ + Get state of hand as well as objects and targets in the scene + """ + qp = self.data.qpos.ravel().copy() + qv = self.data.qvel.ravel().copy() + hand_qpos = qp[:30] + obj_pos = self.data.body_xpos[self.obj_bid].ravel() + palm_pos = self.data.site_xpos[self.S_grasp_sid].ravel() + target_pos = self.data.site_xpos[self.target_obj_sid].ravel() + return dict(hand_qpos=hand_qpos, obj_pos=obj_pos, target_pos=target_pos, palm_pos=palm_pos, + qpos=qp, qvel=qv) + + def set_env_state(self, state_dict): + """ + Set the state which includes hand as well as objects and targets in the scene + """ + qp = state_dict['qpos'] + qv = state_dict['qvel'] + obj_pos = state_dict['obj_pos'] + target_pos = state_dict['target_pos'] + self.set_state(qp, qv) + self.model.body_pos[self.obj_bid] = obj_pos + self.model.site_pos[self.target_obj_sid] = target_pos + self.sim.forward() + + def mj_viewer_setup(self): + self.viewer = MjViewer(self.sim) + self.viewer.cam.azimuth = 90 + self.sim.forward() + self.viewer.cam.distance = 1.5 + + def evaluate_success(self, paths): + num_success = 0 + num_paths = len(paths) + # success if object close to target for 25 steps + for path in paths: + if np.sum(path['env_infos']['goal_achieved']) > 25: + num_success += 1 + success_percentage = num_success*100.0/num_paths + return success_percentage diff --git a/wiserl/env/odrl_envs/adroit/utils/quatmath.py b/wiserl/env/odrl_envs/adroit/utils/quatmath.py new file mode 100644 index 0000000..7fef129 --- /dev/null +++ b/wiserl/env/odrl_envs/adroit/utils/quatmath.py @@ -0,0 +1,164 @@ +import numpy as np +# For testing whether a number is close to zero +_FLOAT_EPS = np.finfo(np.float64).eps +_EPS4 = _FLOAT_EPS * 4.0 + + +def mulQuat(qa, qb): + res = np.zeros(4) + res[0] = qa[0]*qb[0] - qa[1]*qb[1] - qa[2]*qb[2] - qa[3]*qb[3] + res[1] = qa[0]*qb[1] + qa[1]*qb[0] + qa[2]*qb[3] - qa[3]*qb[2] + res[2] = qa[0]*qb[2] - qa[1]*qb[3] + qa[2]*qb[0] + qa[3]*qb[1] + res[3] = qa[0]*qb[3] + qa[1]*qb[2] - qa[2]*qb[1] + qa[3]*qb[0] + return res + +def negQuat(quat): + return np.array([quat[0], -quat[1], -quat[2], -quat[3]]) + +def quat2Vel(quat, dt=1): + axis = quat[1:].copy() + sin_a_2 = np.sqrt(np.sum(axis**2)) + axis = axis/(sin_a_2+1e-8) + speed = 2*np.arctan2(sin_a_2, quat[0])/dt + return speed, axis + +def quatDiff2Vel(quat1, quat2, dt): + neg = negQuat(quat1) + diff = mulQuat(quat2, neg) + return quat2Vel(diff, dt) + + +def axis_angle2quat(axis, angle): + c = np.cos(angle/2) + s = np.sin(angle/2) + return np.array([c, s*axis[0], s*axis[1], s*axis[2]]) + +def euler2mat(euler): + """ Convert Euler Angles to Rotation Matrix. See rotation.py for notes """ + euler = np.asarray(euler, dtype=np.float64) + assert euler.shape[-1] == 3, "Invalid shaped euler {}".format(euler) + + ai, aj, ak = -euler[..., 2], -euler[..., 1], -euler[..., 0] + si, sj, sk = np.sin(ai), np.sin(aj), np.sin(ak) + ci, cj, ck = np.cos(ai), np.cos(aj), np.cos(ak) + cc, cs = ci * ck, ci * sk + sc, ss = si * ck, si * sk + + mat = np.empty(euler.shape[:-1] + (3, 3), dtype=np.float64) + mat[..., 2, 2] = cj * ck + mat[..., 2, 1] = sj * sc - cs + mat[..., 2, 0] = sj * cc + ss + mat[..., 1, 2] = cj * sk + mat[..., 1, 1] = sj * ss + cc + mat[..., 1, 0] = sj * cs - sc + mat[..., 0, 2] = -sj + mat[..., 0, 1] = cj * si + mat[..., 0, 0] = cj * ci + return mat + + +def euler2quat(euler): + """ Convert Euler Angles to Quaternions. See rotation.py for notes """ + euler = np.asarray(euler, dtype=np.float64) + assert euler.shape[-1] == 3, "Invalid shape euler {}".format(euler) + + ai, aj, ak = euler[..., 2] / 2, -euler[..., 1] / 2, euler[..., 0] / 2 + si, sj, sk = np.sin(ai), np.sin(aj), np.sin(ak) + ci, cj, ck = np.cos(ai), np.cos(aj), np.cos(ak) + cc, cs = ci * ck, ci * sk + sc, ss = si * ck, si * sk + + quat = np.empty(euler.shape[:-1] + (4,), dtype=np.float64) + quat[..., 0] = cj * cc + sj * ss + quat[..., 3] = cj * sc - sj * cs + quat[..., 2] = -(cj * ss + sj * cc) + quat[..., 1] = cj * cs - sj * sc + return quat + + +def mat2euler(mat): + """ Convert Rotation Matrix to Euler Angles. See rotation.py for notes """ + mat = np.asarray(mat, dtype=np.float64) + assert mat.shape[-2:] == (3, 3), "Invalid shape matrix {}".format(mat) + + cy = np.sqrt(mat[..., 2, 2] * mat[..., 2, 2] + mat[..., 1, 2] * mat[..., 1, 2]) + condition = cy > _EPS4 + euler = np.empty(mat.shape[:-1], dtype=np.float64) + euler[..., 2] = np.where(condition, + -np.arctan2(mat[..., 0, 1], mat[..., 0, 0]), + -np.arctan2(-mat[..., 1, 0], mat[..., 1, 1])) + euler[..., 1] = np.where(condition, + -np.arctan2(-mat[..., 0, 2], cy), + -np.arctan2(-mat[..., 0, 2], cy)) + euler[..., 0] = np.where(condition, + -np.arctan2(mat[..., 1, 2], mat[..., 2, 2]), + 0.0) + return euler + + +def mat2quat(mat): + """ Convert Rotation Matrix to Quaternion. See rotation.py for notes """ + mat = np.asarray(mat, dtype=np.float64) + assert mat.shape[-2:] == (3, 3), "Invalid shape matrix {}".format(mat) + + Qxx, Qyx, Qzx = mat[..., 0, 0], mat[..., 0, 1], mat[..., 0, 2] + Qxy, Qyy, Qzy = mat[..., 1, 0], mat[..., 1, 1], mat[..., 1, 2] + Qxz, Qyz, Qzz = mat[..., 2, 0], mat[..., 2, 1], mat[..., 2, 2] + # Fill only lower half of symmetric matrix + K = np.zeros(mat.shape[:-2] + (4, 4), dtype=np.float64) + K[..., 0, 0] = Qxx - Qyy - Qzz + K[..., 1, 0] = Qyx + Qxy + K[..., 1, 1] = Qyy - Qxx - Qzz + K[..., 2, 0] = Qzx + Qxz + K[..., 2, 1] = Qzy + Qyz + K[..., 2, 2] = Qzz - Qxx - Qyy + K[..., 3, 0] = Qyz - Qzy + K[..., 3, 1] = Qzx - Qxz + K[..., 3, 2] = Qxy - Qyx + K[..., 3, 3] = Qxx + Qyy + Qzz + K /= 3.0 + # TODO: vectorize this -- probably could be made faster + q = np.empty(K.shape[:-2] + (4,)) + it = np.nditer(q[..., 0], flags=['multi_index']) + while not it.finished: + # Use Hermitian eigenvectors, values for speed + vals, vecs = np.linalg.eigh(K[it.multi_index]) + # Select largest eigenvector, reorder to w,x,y,z quaternion + q[it.multi_index] = vecs[[3, 0, 1, 2], np.argmax(vals)] + # Prefer quaternion with positive w + # (q * -1 corresponds to same rotation as q) + if q[it.multi_index][0] < 0: + q[it.multi_index] *= -1 + it.iternext() + return q + + +def quat2euler(quat): + """ Convert Quaternion to Euler Angles. See rotation.py for notes """ + return mat2euler(quat2mat(quat)) + + +def quat2mat(quat): + """ Convert Quaternion to Euler Angles. See rotation.py for notes """ + quat = np.asarray(quat, dtype=np.float64) + assert quat.shape[-1] == 4, "Invalid shape quat {}".format(quat) + + w, x, y, z = quat[..., 0], quat[..., 1], quat[..., 2], quat[..., 3] + Nq = np.sum(quat * quat, axis=-1) + s = 2.0 / Nq + X, Y, Z = x * s, y * s, z * s + wX, wY, wZ = w * X, w * Y, w * Z + xX, xY, xZ = x * X, x * Y, x * Z + yY, yZ, zZ = y * Y, y * Z, z * Z + + mat = np.empty(quat.shape[:-1] + (3, 3), dtype=np.float64) + mat[..., 0, 0] = 1.0 - (yY + zZ) + mat[..., 0, 1] = xY - wZ + mat[..., 0, 2] = xZ + wY + mat[..., 1, 0] = xY + wZ + mat[..., 1, 1] = 1.0 - (xX + zZ) + mat[..., 1, 2] = yZ - wX + mat[..., 2, 0] = xZ - wY + mat[..., 2, 1] = yZ + wX + mat[..., 2, 2] = 1.0 - (xX + yY) + return np.where((Nq > _FLOAT_EPS)[..., np.newaxis, np.newaxis], mat, np.eye(3)) \ No newline at end of file diff --git a/wiserl/env/odrl_envs/antmaze/__init__.py b/wiserl/env/odrl_envs/antmaze/__init__.py new file mode 100644 index 0000000..c743e62 --- /dev/null +++ b/wiserl/env/odrl_envs/antmaze/__init__.py @@ -0,0 +1,517 @@ +import gym +from gym.envs.registration import register +from pathlib import Path +import sys +sys.path.append(str(Path(__file__).parent.absolute())) +from ant import make_ant_maze_env +import ant + +RESET = R = 'r' # Reset position. +GOAL = G = 'g' + +# small mazes +register( + id='antmaze-small-0-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=700, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1], + [1, R, 0, 0, 1], + [1, 1, 1, 0, 1], + [1, G, 0, 0, 1], + [1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) +register( + id='antmaze-small-empty-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=700, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1], + [1, R, 0, 0, 1], + [1, 0, 0, 0, 1], + [1, 0, 0, G, 1], + [1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) +register( + id='antmaze-small-centerblock-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=700, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1], + [1, R, 0, 0, 1], + [1, 0, 1, 0, 1], + [1, 0, 0, G, 1], + [1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) +register( + id='antmaze-small-lshape-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=700, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1], + [1, R, 1, 1, 1], + [1, 0, 1, 1, 1], + [1, 0, 0, G, 1], + [1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) + +register( + id='antmaze-small-zshape-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=700, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1], + [1, R, 0, 1, 1], + [1, 1, 0, 1, 1], + [1, 1, 0, G, 1], + [1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) + +register( + id='antmaze-small-reversel-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=700, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1], + [1, R, 0, 0, 1], + [1, 1, 1, 0, 1], + [1, 1, 1, G, 1], + [1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) + +register( + id='antmaze-small-reverseu-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=700, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1], + [1, R, 0, 0, 1], + [1, 0, 1, 0, 1], + [1, 0, 1, G, 1], + [1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) + +# medium mazes +register( + id='antmaze-medium-0-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=1000, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 1, 1, 0, 0, 1], + [1, 0, 0, 1, 0, 0, 0, 1], + [1, 1, 0, 0, 0, 1, 1, 1], + [1, 0, 0, 1, 0, 0, 0, 1], + [1, 0, 1, 0, 0, 1, 0, 1], + [1, 0, 0, 0, 1, 0, G, 1], + [1, 1, 1, 1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) +register( + id='antmaze-medium-1-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=1000, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 1, 0, 0, 0, 1], + [1, 0, 0, 1, 0, 1, 1, 1], + [1, 1, 0, 0, 0, 0, 0, 1], + [1, 0, 0, 1, 1, 0, 0, 1], + [1, 0, 1, 0, 1, 1, 0, 1], + [1, 0, 0, 0, 1, 0, G, 1], + [1, 1, 1, 1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) +register( + id='antmaze-medium-2-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=1000, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 0, 1, 1, 0, 1], + [1, 0, 1, 0, 1, 0, 0, 1], + [1, 0, 0, 0, 0, 0, 0, 1], + [1, 0, 0, 1, 0, 1, 0, 1], + [1, 0, 1, 1, 0, 1, 0, 1], + [1, 0, 1, 0, 0, 0, G, 1], + [1, 1, 1, 1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) +register( + id='antmaze-medium-3-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=1000, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 1, 0, 0, 0, 1], + [1, 0, 0, 1, 1, 1, 0, 1], + [1, 1, 0, 0, 0, 0, 0, 1], + [1, 0, 0, 1, 1, 1, 0, 1], + [1, 1, 0, 1, 0, 0, 0, 1], + [1, 0, 0, 0, 0, 0, G, 1], + [1, 1, 1, 1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) + +register( + id='antmaze-medium-4-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=1000, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 1, 0, 1, 0, 1], + [1, 0, 0, 0, 0, 0, 0, 1], + [1, 1, 0, 1, 1, 0, 0, 1], + [1, 0, 0, 0, 1, 1, 0, 1], + [1, 0, 0, 1, 0, 0, 0, 1], + [1, 0, 0, 0, 0, 0, G, 1], + [1, 1, 1, 1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) + +register( + id='antmaze-medium-5-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=1000, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 1, 1, 0, 0, 1], + [1, 0, 0, 0, 0, 0, 0, 1], + [1, 1, 0, 0, 0, 1, 0, 1], + [1, 0, 0, 1, 0, 0, 0, 1], + [1, 0, 1, 0, 0, 1, 0, 1], + [1, 0, 0, 0, 1, 0, G, 1], + [1, 1, 1, 1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) + +register( + id='antmaze-medium-6-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=1000, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 0, 0, 1, 0, 1], + [1, 1, 1, 0, 0, 0, 0, 1], + [1, 1, 0, 0, 0, 1, 1, 1], + [1, 0, 1, 1, 0, 0, 0, 1], + [1, 0, 0, 0, 0, 0, 0, 1], + [1, 1, 1, 0, 1, 0, G, 1], + [1, 1, 1, 1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) + +# large mazes +register( + id='antmaze-large-0-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=1000, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 0, 0, 1, 0, 0, 0, 0, 0, 1], + [1, 0, 1, 1, 0, 1, 0, 1, 0, 1, 0, 1], + [1, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 1], + [1, 0, 1, 1, 1, 1, 0, 1, 1, 1, 0, 1], + [1, 0, 0, 1, 0, 1, 0, 0, 0, 0, 0, 1], + [1, 1, 0, 1, 0, 1, 0, 1, 0, 1, 1, 1], + [1, 0, 0, 1, 0, 0, 0, 1, 0, 0, G, 1], + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) + +register( + id='antmaze-large-1-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=1000, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 1, 0, 0, 1, 1, 0, 0, 0, 0, 1], + [1, 0, 1, 1, 0, 1, 0, 1, 0, 1, 0, 1], + [1, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 1], + [1, 0, 1, 1, 1, 1, 0, 0, 0, 1, 0, 1], + [1, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1], + [1, 1, 0, 1, 0, 1, 0, 1, 0, 1, 1, 1], + [1, 0, 0, 1, 0, 0, 0, 1, 0, 0, G, 1], + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) + +register( + id='antmaze-large-2-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=1000, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 0, 0, 0, 0, 0, 0, 1, 0, 1], + [1, 0, 0, 1, 1, 1, 0, 1, 0, 1, 0, 1], + [1, 0, 0, 0, 0, 1, 0, 1, 0, 0, 0, 1], + [1, 0, 1, 1, 1, 0, 0, 1, 1, 1, 0, 1], + [1, 0, 0, 1, 0, 1, 0, 0, 0, 0, 0, 1], + [1, 1, 0, 1, 0, 1, 0, 1, 0, 1, 1, 1], + [1, 0, 0, 1, 0, 0, 0, 1, 0, 0, G, 1], + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) +register( + id='antmaze-large-3-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=1000, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 0, 0, 0, 0, 1, 0, 0, 0, 1], + [1, 0, 1, 1, 0, 1, 0, 1, 0, 1, 0, 1], + [1, 0, 1, 0, 0, 0, 0, 0, 0, 1, 0, 1], + [1, 0, 0, 0, 1, 1, 0, 0, 0, 1, 0, 1], + [1, 0, 0, 1, 0, 1, 0, 1, 0, 0, 0, 1], + [1, 1, 0, 0, 0, 1, 0, 1, 1, 0, 0, 1], + [1, 0, 0, 1, 0, 0, 0, 1, 0, 0, G, 1], + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) + +register( + id='antmaze-large-4-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=1000, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1], + [1, 0, 0, 1, 0, 1, 0, 0, 0, 1, 0, 1], + [1, 1, 0, 0, 0, 1, 0, 1, 0, 1, 0, 1], + [1, 0, 1, 0, 1, 1, 0, 1, 0, 1, 1, 1], + [1, 0, 0, 0, 0, 1, 0, 1, 0, 0, 0, 1], + [1, 0, 1, 1, 0, 0, 0, 1, 0, 0, 0, 1], + [1, 0, 0, 1, 0, 0, 0, 1, 0, 0, G, 1], + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) + +register( + id='antmaze-large-5-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=1000, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1], + [1, 0, 1, 0, 1, 1, 0, 1, 1, 1, 0, 1], + [1, 0, 0, 0, 0, 1, 0, 1, 0, 0, 0, 1], + [1, 1, 1, 0, 1, 1, 0, 1, 0, 1, 0, 1], + [1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1], + [1, 1, 0, 1, 0, 1, 0, 1, 1, 0, 0, 1], + [1, 0, 0, 1, 0, 1, 0, 1, 0, 0, G, 1], + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) + +register( + id='antmaze-large-6-v0', + entry_point='ant:make_ant_maze_env', + max_episode_steps=1000, + kwargs={ + 'maze_map': [ + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 0, 0, 0, 1, 0, 0, 0, 0, 1], + [1, 0, 0, 0, 1, 0, 1, 1, 1, 1, 0, 1], + [1, 1, 1, 0, 1, 0, 1, 0, 0, 0, 0, 1], + [1, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1], + [1, 1, 1, 1, 0, 0, 0, 0, 0, 0, 1, 1], + [1, 0, 0, 0, 0, 1, 1, 1, 0, 0, 0, 1], + [1, 0, 1, 1, 0, 1, 0, 0, 0, 0, G, 1], + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1] + ], + 'reward_type':'sparse', + 'non_zero_reset':False, + 'eval':True, + 'maze_size_scaling': 4.0, + 'ref_min_score': 0.0, + 'ref_max_score': 1.0, + 'v2_resets': True, + } +) \ No newline at end of file diff --git a/wiserl/env/odrl_envs/antmaze/ant.py b/wiserl/env/odrl_envs/antmaze/ant.py new file mode 100644 index 0000000..0fb9ce8 --- /dev/null +++ b/wiserl/env/odrl_envs/antmaze/ant.py @@ -0,0 +1,213 @@ +# Copyright 2018 The TensorFlow Authors All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Wrapper for creating the ant environment.""" + +import numpy as np +import mujoco_py +import os + +from gym import utils +from gym.envs.mujoco import mujoco_env + +import mujoco_goal_env +import goal_reaching_env +import maze_env +import wrappers + + +GYM_ASSETS_DIR = os.path.join( + os.path.dirname(mujoco_goal_env.__file__), + 'assets') + +class AntEnv(mujoco_env.MujocoEnv, utils.EzPickle): + """Basic ant locomotion environment.""" + FILE = os.path.join(GYM_ASSETS_DIR, 'ant.xml') + + def __init__(self, file_path=None, expose_all_qpos=False, + expose_body_coms=None, expose_body_comvels=None, non_zero_reset=False): + if file_path is None: + file_path = self.FILE + + self._expose_all_qpos = expose_all_qpos + self._expose_body_coms = expose_body_coms + self._expose_body_comvels = expose_body_comvels + self._body_com_indices = {} + self._body_comvel_indices = {} + + self._non_zero_reset = non_zero_reset + + mujoco_env.MujocoEnv.__init__(self, file_path, 5) + utils.EzPickle.__init__(self) + + @property + def physics(self): + # Check mujoco version is greater than version 1.50 to call correct physics + # model containing PyMjData object for getting and setting position/velocity. + # Check https://github.com/openai/mujoco-py/issues/80 for updates to api. + if mujoco_py.get_version() >= '1.50': + return self.sim + else: + return self.model + + def _step(self, a): + return self.step(a) + + def step(self, a): + xposbefore = self.get_body_com("torso")[0] + self.do_simulation(a, self.frame_skip) + xposafter = self.get_body_com("torso")[0] + forward_reward = (xposafter - xposbefore) / self.dt + ctrl_cost = .5 * np.square(a).sum() + contact_cost = 0.5 * 1e-3 * np.sum( + np.square(np.clip(self.sim.data.cfrc_ext, -1, 1))) + survive_reward = 1.0 + reward = forward_reward - ctrl_cost - contact_cost + survive_reward + state = self.state_vector() + notdone = np.isfinite(state).all() \ + and state[2] >= 0.2 and state[2] <= 1.0 + done = not notdone + ob = self._get_obs() + return ob, reward, done, dict( + reward_forward=forward_reward, + reward_ctrl=-ctrl_cost, + reward_contact=-contact_cost, + reward_survive=survive_reward) + + def _get_obs(self): + # No cfrc observation. + if self._expose_all_qpos: + obs = np.concatenate([ + self.physics.data.qpos.flat[:15], # Ensures only ant obs. + self.physics.data.qvel.flat[:14], + ]) + else: + obs = np.concatenate([ + self.physics.data.qpos.flat[2:15], + self.physics.data.qvel.flat[:14], + ]) + + if self._expose_body_coms is not None: + for name in self._expose_body_coms: + com = self.get_body_com(name) + if name not in self._body_com_indices: + indices = range(len(obs), len(obs) + len(com)) + self._body_com_indices[name] = indices + obs = np.concatenate([obs, com]) + + if self._expose_body_comvels is not None: + for name in self._expose_body_comvels: + comvel = self.get_body_comvel(name) + if name not in self._body_comvel_indices: + indices = range(len(obs), len(obs) + len(comvel)) + self._body_comvel_indices[name] = indices + obs = np.concatenate([obs, comvel]) + return obs + + def reset_model(self): + qpos = self.init_qpos + self.np_random.uniform( + size=self.model.nq, low=-.1, high=.1) + qvel = self.init_qvel + self.np_random.randn(self.model.nv) * .1 + + if self._non_zero_reset: + """Now the reset is supposed to be to a non-zero location""" + reset_location = self._get_reset_location() + qpos[:2] = reset_location + + # Set everything other than ant to original position and 0 velocity. + qpos[15:] = self.init_qpos[15:] + qvel[14:] = 0. + self.set_state(qpos, qvel) + return self._get_obs() + + def viewer_setup(self): + self.viewer.cam.distance = self.model.stat.extent * 1.6 + self.viewer.cam.elevation = -60 + # self.viewer.cam.lookat = np.array([0.5, 0.5, 1.]) + + def get_xy(self): + return self.physics.data.qpos[:2] + + def set_xy(self, xy): + qpos = np.copy(self.physics.data.qpos) + qpos[0] = xy[0] + qpos[1] = xy[1] + qvel = self.physics.data.qvel + self.set_state(qpos, qvel) + + +class GoalReachingAntEnv(goal_reaching_env.GoalReachingEnv, AntEnv): + """Ant locomotion rewarded for goal-reaching.""" + BASE_ENV = AntEnv + + def __init__(self, goal_sampler=goal_reaching_env.disk_goal_sampler, + file_path=None, + expose_all_qpos=False, non_zero_reset=False, eval=False, reward_type='dense', **kwargs): + goal_reaching_env.GoalReachingEnv.__init__(self, goal_sampler, eval=eval, reward_type=reward_type) + AntEnv.__init__(self, + file_path=file_path, + expose_all_qpos=expose_all_qpos, + expose_body_coms=None, + expose_body_comvels=None, + non_zero_reset=non_zero_reset) + +class AntMazeEnv(maze_env.MazeEnv, GoalReachingAntEnv): + """Ant navigating a maze.""" + LOCOMOTION_ENV = GoalReachingAntEnv + + def __init__(self, goal_sampler=None, expose_all_qpos=True, + reward_type='dense', v2_resets=False, + *args, **kwargs): + if goal_sampler is None: + goal_sampler = lambda np_rand: maze_env.MazeEnv.goal_sampler(self, np_rand) + maze_env.MazeEnv.__init__( + self, *args, manual_collision=False, + goal_sampler=goal_sampler, + expose_all_qpos=expose_all_qpos, + reward_type=reward_type, + **kwargs) + + ## We set the target foal here for evaluation + self.set_target() + self.v2_resets = v2_resets + + def reset(self): + if self.v2_resets: + """ + The target goal for evaluation in antmazes is randomized. + antmazes-v0 and -v1 resulted in really high-variance evaluations + because the target goal was set once at the seed level. This led to + each run running evaluations with one particular goal. To accurately + cover each goal, this requires about 50-100 seeds, which might be + computationally infeasible. As an alternate fix, to reduce variance + in result reporting, we are creating the v2 environments + which use the same offline dataset as v0 environments, with the distinction + that the randomization of goals during evaluation is performed at the level of + each rollout. Thus running a few seeds, but performing the final evaluation + over 100-200 episodes will give a valid estimate of an algorithm's performance. + """ + self.set_target() + return super().reset() + + def set_target(self, target_location=None): + return self.set_target_goal(target_location) + + def seed(self, seed=0): + mujoco_env.MujocoEnv.seed(self, seed) + +def make_ant_maze_env(**kwargs): + env = AntMazeEnv(**kwargs) + return wrappers.NormalizedBoxEnv(env) + diff --git a/wiserl/env/odrl_envs/antmaze/assets/ant.xml b/wiserl/env/odrl_envs/antmaze/assets/ant.xml new file mode 100644 index 0000000..5a510be --- /dev/null +++ b/wiserl/env/odrl_envs/antmaze/assets/ant.xml @@ -0,0 +1,81 @@ + + + diff --git a/wiserl/env/odrl_envs/antmaze/assets/point.xml b/wiserl/env/odrl_envs/antmaze/assets/point.xml new file mode 100644 index 0000000..f8181bb --- /dev/null +++ b/wiserl/env/odrl_envs/antmaze/assets/point.xml @@ -0,0 +1,30 @@ + + + diff --git a/wiserl/env/odrl_envs/antmaze/call_antmaze_env.py b/wiserl/env/odrl_envs/antmaze/call_antmaze_env.py new file mode 100644 index 0000000..396b9c9 --- /dev/null +++ b/wiserl/env/odrl_envs/antmaze/call_antmaze_env.py @@ -0,0 +1,42 @@ +from typing import Dict +import gym +import d4rl + + +def call_antmaze_env(env_config: Dict) -> gym.Env: + env_name = env_config['env_name'].lower() # eg. "antmaze_small_lshape" + if '_' in env_name: + env_name = env_name.replace('_', '-') + # decide which task it is, support the following tasks + # antmaze - small - empty + # - lshape + # - centerblock + # - brokenjoint + # - reversel + # - reverseu + # - zshape + # - medium - 1/2/3/4/5/6 + # - large - 1/2/3/4/5/6 + assert env_name.startswith('antmaze') + assert any([size in env_name for size in ['small', 'medium', 'large']]) + + shift_level = env_config['shift_level'] + + if shift_level is None: + if 'small' in env_name: + return gym.make('antmaze-umaze-v0') + elif 'medium' in env_name: + return gym.make('antmaze-medium-0-v0') + else: + return gym.make('antmaze-large-0-v0') + else: + if 'small' in env_name: + assert any([size in shift_level for size in ['empty', 'lshape', 'centerblock', 'reversel', 'reverseu', 'zshape']]) + env_name += '-' + str(shift_level) + '-v0' + elif 'medium' in env_name: + assert any([size in shift_level for size in ['0','1', '2', '3', '4', '5', '6']]) + env_name += '-' + str(shift_level) + '-v0' + else: + assert any([size in shift_level for size in ['1', '2', '3', '4', '5', '6']]) + env_name += '-' + str(shift_level) + '-v0' + return gym.make(env_name) \ No newline at end of file diff --git a/wiserl/env/odrl_envs/antmaze/common.py b/wiserl/env/odrl_envs/antmaze/common.py new file mode 100644 index 0000000..56c415c --- /dev/null +++ b/wiserl/env/odrl_envs/antmaze/common.py @@ -0,0 +1,21 @@ + + +def run_policy_on_env(policy_fn, env, truncate_episode_at=None, + first_obs=None): + if first_obs is None: + obs = env.reset() + else: + obs = first_obs + + trajectory = [] + step_num = 0 + while True: + act = policy_fn(obs) + next_obs, rew, done, _ = env.step(act) + trajectory.append((obs, act, rew, done)) + obs = next_obs + step_num += 1 + if (done or + (truncate_episode_at is not None and step_num >= truncate_episode_at)): + break + return trajectory diff --git a/wiserl/env/odrl_envs/antmaze/goal_reaching_env.py b/wiserl/env/odrl_envs/antmaze/goal_reaching_env.py new file mode 100644 index 0000000..eef1569 --- /dev/null +++ b/wiserl/env/odrl_envs/antmaze/goal_reaching_env.py @@ -0,0 +1,58 @@ +import numpy as np + + +def disk_goal_sampler(np_random, goal_region_radius=10.): + th = 2 * np.pi * np_random.uniform() + radius = goal_region_radius * np_random.uniform() + return radius * np.array([np.cos(th), np.sin(th)]) + +def constant_goal_sampler(np_random, location=10.0 * np.ones([2])): + return location + +class GoalReachingEnv(object): + """General goal-reaching environment.""" + BASE_ENV = None # Must be specified by child class. + + def __init__(self, goal_sampler, eval=False, reward_type='dense'): + self._goal_sampler = goal_sampler + self._goal = np.ones([2]) + self.target_goal = self._goal + + # This flag is used to make sure that when using this environment + # for evaluation, that is no goals are appended to the state + self.eval = eval + + # This is the reward type fed as input to the goal confitioned policy + self.reward_type = reward_type + + def _get_obs(self): + base_obs = self.BASE_ENV._get_obs(self) + goal_direction = self._goal - self.get_xy() + if not self.eval: + obs = np.concatenate([base_obs, goal_direction]) + return obs + else: + return base_obs + + def step(self, a): + self.BASE_ENV.step(self, a) + if self.reward_type == 'dense': + reward = -np.linalg.norm(self.target_goal - self.get_xy()) + elif self.reward_type == 'sparse': + reward = 1.0 if np.linalg.norm(self.get_xy() - self.target_goal) <= 0.5 else 0.0 + + done = False + # Terminate episode when we reach a goal + if self.eval and np.linalg.norm(self.get_xy() - self.target_goal) <= 0.5: + done = True + + obs = self._get_obs() + return obs, reward, done, {} + + def reset_model(self): + if self.target_goal is not None or self.eval: + self._goal = self.target_goal + else: + self._goal = self._goal_sampler(self.np_random) + + return self.BASE_ENV.reset_model(self) \ No newline at end of file diff --git a/wiserl/env/odrl_envs/antmaze/maze_env.py b/wiserl/env/odrl_envs/antmaze/maze_env.py new file mode 100644 index 0000000..791ac13 --- /dev/null +++ b/wiserl/env/odrl_envs/antmaze/maze_env.py @@ -0,0 +1,377 @@ +# Copyright 2018 The TensorFlow Authors All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Adapted from efficient-hrl maze_env.py.""" + +import os +import tempfile +import xml.etree.ElementTree as ET +import math +import numpy as np +import gym +from copy import deepcopy + +RESET = R = 'r' # Reset position. +GOAL = G = 'g' + +# Maze specifications for dataset generation +U_MAZE = [[1, 1, 1, 1, 1], + [1, R, 0, 0, 1], + [1, 1, 1, 0, 1], + [1, G, 0, 0, 1], + [1, 1, 1, 1, 1]] + +BIG_MAZE = [[1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 1, 1, 0, 0, 1], + [1, 0, 0, 1, 0, 0, G, 1], + [1, 1, 0, 0, 0, 1, 1, 1], + [1, 0, 0, 1, 0, 0, 0, 1], + [1, G, 1, 0, 0, 1, 0, 1], + [1, 0, 0, 0, 1, G, 0, 1], + [1, 1, 1, 1, 1, 1, 1, 1]] + +HARDEST_MAZE = [[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 0, 0, 1, G, 0, 0, 0, 0, 1], + [1, 0, 1, 1, 0, 1, 0, 1, 0, 1, 0, 1], + [1, 0, 0, 0, 0, G, 0, 1, 0, 0, G, 1], + [1, 0, 1, 1, 1, 1, 0, 1, 1, 1, 0, 1], + [1, 0, G, 1, 0, 1, 0, 0, 0, 0, 0, 1], + [1, 1, 0, 1, 0, 1, 0, 1, 0, 1, 1, 1], + [1, 0, 0, 1, G, 0, G, 1, 0, G, 0, 1], + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]] + +# Maze specifications with a single target goal +U_MAZE_TEST = [[1, 1, 1, 1, 1], + [1, R, 0, 0, 1], + [1, 1, 1, 0, 1], + [1, G, 0, 0, 1], + [1, 1, 1, 1, 1]] + +BIG_MAZE_TEST = [[1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 1, 1, 0, 0, 1], + [1, 0, 0, 1, 0, 0, 0, 1], + [1, 1, 0, 0, 0, 1, 1, 1], + [1, 0, 0, 1, 0, 0, 0, 1], + [1, 0, 1, 0, 0, 1, 0, 1], + [1, 0, 0, 0, 1, 0, G, 1], + [1, 1, 1, 1, 1, 1, 1, 1]] + +HARDEST_MAZE_TEST = [[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 0, 0, 1, 0, 0, 0, 0, 0, 1], + [1, 0, 1, 1, 0, 1, 0, 1, 0, 1, 0, 1], + [1, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 1], + [1, 0, 1, 1, 1, 1, 0, 1, 1, 1, 0, 1], + [1, 0, 0, 1, 0, 1, 0, 0, 0, 0, 0, 1], + [1, 1, 0, 1, 0, 1, 0, 1, 0, 1, 1, 1], + [1, 0, 0, 1, 0, 0, 0, 1, 0, G, 0, 1], + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]] + +# Maze specifications for evaluation +U_MAZE_EVAL = [[1, 1, 1, 1, 1], + [1, 0, 0, R, 1], + [1, 0, 1, 1, 1], + [1, 0, 0, G, 1], + [1, 1, 1, 1, 1]] + +BIG_MAZE_EVAL = [[1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 0, 0, 0, G, 1], + [1, 0, 1, 0, 1, 1, 0, 1], + [1, 0, 0, 0, 0, 1, 0, 1], + [1, 1, 1, 0, 0, 1, 1, 1], + [1, G, 0, 0, 0, 0, 0, 1], + [1, 0, 0, 1, 1, G, 0, 1], + [1, 1, 1, 1, 1, 1, 1, 1]] + +HARDEST_MAZE_EVAL = [[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 1, G, 0, 0, 1, 0, G, 0, 1], + [1, 1, 0, 1, 1, 1, 0, 1, 0, 1, 0, 1], + [1, 0, 0, 1, 0, 1, G, 0, 0, 0, 0, 1], + [1, 0, 1, 1, 0, 1, 0, 0, 1, 1, 0, 1], + [1, G, 0, 0, 0, 0, 0, 1, 0, 0, 0, 1], + [1, 0, 1, 1, 0, 1, 0, 1, 0, 1, 1, 1], + [1, 0, 0, 0, G, 1, G, 0, 0, 0, G, 1], + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]] + +U_MAZE_EVAL_TEST = [[1, 1, 1, 1, 1], + [1, 0, 0, R, 1], + [1, 0, 1, 1, 1], + [1, 0, 0, G, 1], + [1, 1, 1, 1, 1]] + +BIG_MAZE_EVAL_TEST = [[1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 0, 0, 0, G, 1], + [1, 0, 1, 0, 1, 1, 0, 1], + [1, 0, 0, 0, 0, 1, 0, 1], + [1, 1, 1, 0, 0, 1, 1, 1], + [1, 0, 0, 0, 0, 0, 0, 1], + [1, 0, 0, 1, 1, 0, 0, 1], + [1, 1, 1, 1, 1, 1, 1, 1]] + +HARDEST_MAZE_EVAL_TEST = [[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1], + [1, R, 0, 1, 0, 0, 0, 1, 0, G, 0, 1], + [1, 1, 0, 1, 1, 1, 0, 1, 0, 1, 0, 1], + [1, 0, 0, 1, 0, 1, 0, 0, 0, 0, 0, 1], + [1, 0, 1, 1, 0, 1, 0, 0, 1, 1, 0, 1], + [1, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 1], + [1, 0, 1, 1, 0, 1, 0, 1, 0, 1, 1, 1], + [1, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 1], + [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]] + + +class MazeEnv(gym.Env): + LOCOMOTION_ENV = None # Must be specified by child class. + + def __init__( + self, + maze_map, + maze_size_scaling, + maze_height=0.5, + manual_collision=False, + non_zero_reset=False, + reward_type='dense', + *args, + **kwargs): + if self.LOCOMOTION_ENV is None: + raise ValueError('LOCOMOTION_ENV is unspecified.') + + xml_path = self.LOCOMOTION_ENV.FILE + tree = ET.parse(xml_path) + worldbody = tree.find(".//worldbody") + + self._maze_map = maze_map + + self._maze_height = maze_height + self._maze_size_scaling = maze_size_scaling + self._manual_collision = manual_collision + + self._maze_map = maze_map + + # Obtain a numpy array form for a maze map in case we want to reset + # to multiple starting states + temp_maze_map = deepcopy(self._maze_map) + for i in range(len(maze_map)): + for j in range(len(maze_map[0])): + if temp_maze_map[i][j] in [RESET,]: + temp_maze_map[i][j] = 0 + elif temp_maze_map[i][j] in [GOAL,]: + temp_maze_map[i][j] = 1 + + self._np_maze_map = np.array(temp_maze_map) + + torso_x, torso_y = self._find_robot() + self._init_torso_x = torso_x + self._init_torso_y = torso_y + + for i in range(len(self._maze_map)): + for j in range(len(self._maze_map[0])): + struct = self._maze_map[i][j] + if struct == 1: # Unmovable block. + # Offset all coordinates so that robot starts at the origin. + ET.SubElement( + worldbody, "geom", + name="block_%d_%d" % (i, j), + pos="%f %f %f" % (j * self._maze_size_scaling - torso_x, + i * self._maze_size_scaling - torso_y, + self._maze_height / 2 * self._maze_size_scaling), + size="%f %f %f" % (0.5 * self._maze_size_scaling, + 0.5 * self._maze_size_scaling, + self._maze_height / 2 * self._maze_size_scaling), + type="box", + material="", + contype="1", + conaffinity="1", + rgba="0.7 0.5 0.3 1.0", + ) + + torso = tree.find(".//body[@name='torso']") + geoms = torso.findall(".//geom") + + _, file_path = tempfile.mkstemp(text=True, suffix='.xml') + tree.write(file_path) + + self.LOCOMOTION_ENV.__init__(self, *args, file_path=file_path, non_zero_reset=non_zero_reset, reward_type=reward_type, **kwargs) + + self.target_goal = None + + def _xy_to_rowcol(self, xy): + size_scaling = self._maze_size_scaling + xy = (max(xy[0], 1e-4), max(xy[1], 1e-4)) + return (int(1 + (xy[1]) / size_scaling), + int(1 + (xy[0]) / size_scaling)) + + def _get_reset_location(self,): + prob = (1.0 - self._np_maze_map) / np.sum(1.0 - self._np_maze_map) + prob_row = np.sum(prob, 1) + row_sample = np.random.choice(np.arange(self._np_maze_map.shape[0]), p=prob_row) + col_sample = np.random.choice(np.arange(self._np_maze_map.shape[1]), p=prob[row_sample] * 1.0 / prob_row[row_sample]) + reset_location = self._rowcol_to_xy((row_sample, col_sample)) + + # Add some random noise + random_x = np.random.uniform(low=0, high=0.5) * 0.5 * self._maze_size_scaling + random_y = np.random.uniform(low=0, high=0.5) * 0.5 * self._maze_size_scaling + + return (max(reset_location[0] + random_x, 0), max(reset_location[1] + random_y, 0)) + + def _rowcol_to_xy(self, rowcol, add_random_noise=False): + row, col = rowcol + x = col * self._maze_size_scaling - self._init_torso_x + y = row * self._maze_size_scaling - self._init_torso_y + if add_random_noise: + x = x + np.random.uniform(low=0, high=self._maze_size_scaling * 0.25) + y = y + np.random.uniform(low=0, high=self._maze_size_scaling * 0.25) + return (x, y) + + def goal_sampler(self, np_random, only_free_cells=True, interpolate=True): + valid_cells = [] + goal_cells = [] + + for i in range(len(self._maze_map)): + for j in range(len(self._maze_map[0])): + if self._maze_map[i][j] in [0, RESET, GOAL] or not only_free_cells: + valid_cells.append((i, j)) + if self._maze_map[i][j] == GOAL: + goal_cells.append((i, j)) + + # If there is a 'goal' designated, use that. Otherwise, any valid cell can + # be a goal. + sample_choices = goal_cells if goal_cells else valid_cells + cell = sample_choices[np_random.choice(len(sample_choices))] + xy = self._rowcol_to_xy(cell, add_random_noise=True) + + random_x = np.random.uniform(low=0, high=0.5) * 0.25 * self._maze_size_scaling + random_y = np.random.uniform(low=0, high=0.5) * 0.25 * self._maze_size_scaling + + xy = (max(xy[0] + random_x, 0), max(xy[1] + random_y, 0)) + + return xy + + def set_target_goal(self, goal_input=None): + if goal_input is None: + self.target_goal = self.goal_sampler(np.random) + else: + self.target_goal = goal_input + + # print ('Target Goal: ', self.target_goal) + ## Make sure that the goal used in self._goal is also reset: + self._goal = self.target_goal + + def _find_robot(self): + structure = self._maze_map + size_scaling = self._maze_size_scaling + for i in range(len(structure)): + for j in range(len(structure[0])): + if structure[i][j] == RESET: + return j * size_scaling, i * size_scaling + raise ValueError('No robot in maze specification.') + + def _is_in_collision(self, pos): + x, y = pos + structure = self._maze_map + size_scaling = self._maze_size_scaling + for i in range(len(structure)): + for j in range(len(structure[0])): + if structure[i][j] == 1: + minx = j * size_scaling - size_scaling * 0.5 - self._init_torso_x + maxx = j * size_scaling + size_scaling * 0.5 - self._init_torso_x + miny = i * size_scaling - size_scaling * 0.5 - self._init_torso_y + maxy = i * size_scaling + size_scaling * 0.5 - self._init_torso_y + if minx <= x <= maxx and miny <= y <= maxy: + return True + return False + + def step(self, action): + if self._manual_collision: + old_pos = self.get_xy() + inner_next_obs, inner_reward, done, info = self.LOCOMOTION_ENV.step(self, action) + new_pos = self.get_xy() + if self._is_in_collision(new_pos): + self.set_xy(old_pos) + else: + inner_next_obs, inner_reward, done, info = self.LOCOMOTION_ENV.step(self, action) + next_obs = self._get_obs() + return next_obs, inner_reward, done, info + + def _get_best_next_rowcol(self, current_rowcol, target_rowcol): + """Runs BFS to find shortest path to target and returns best next rowcol. + Add obstacle avoidance""" + current_rowcol = tuple(current_rowcol) + target_rowcol = tuple(target_rowcol) + if target_rowcol == current_rowcol: + return target_rowcol + + visited = {} + to_visit = [target_rowcol] + while to_visit: + next_visit = [] + for rowcol in to_visit: + visited[rowcol] = True + row, col = rowcol + left = (row, col - 1) + right = (row, col + 1) + down = (row + 1, col) + up = (row - 1, col) + for next_rowcol in [left, right, down, up]: + if next_rowcol == current_rowcol: # Found a shortest path. + return rowcol + next_row, next_col = next_rowcol + if next_row < 0 or next_row >= len(self._maze_map): + continue + if next_col < 0 or next_col >= len(self._maze_map[0]): + continue + if self._maze_map[next_row][next_col] not in [0, RESET, GOAL]: + continue + if next_rowcol in visited: + continue + next_visit.append(next_rowcol) + to_visit = next_visit + + raise ValueError('No path found to target.') + + def create_navigation_policy(self, + goal_reaching_policy_fn, + obs_to_robot=lambda obs: obs[:2], + obs_to_target=lambda obs: obs[-2:], + relative=False): + """Creates a navigation policy by guiding a sub-policy to waypoints.""" + + def policy_fn(obs): + # import ipdb; ipdb.set_trace() + robot_x, robot_y = obs_to_robot(obs) + robot_row, robot_col = self._xy_to_rowcol([robot_x, robot_y]) + target_x, target_y = self.target_goal + if relative: + target_x += robot_x # Target is given in relative coordinates. + target_y += robot_y + target_row, target_col = self._xy_to_rowcol([target_x, target_y]) + print ('Target: ', target_row, target_col, target_x, target_y) + print ('Robot: ', robot_row, robot_col, robot_x, robot_y) + + waypoint_row, waypoint_col = self._get_best_next_rowcol( + [robot_row, robot_col], [target_row, target_col]) + + if waypoint_row == target_row and waypoint_col == target_col: + waypoint_x = target_x + waypoint_y = target_y + else: + waypoint_x, waypoint_y = self._rowcol_to_xy([waypoint_row, waypoint_col], add_random_noise=True) + + goal_x = waypoint_x - robot_x + goal_y = waypoint_y - robot_y + + print ('Waypoint: ', waypoint_row, waypoint_col, waypoint_x, waypoint_y) + + return goal_reaching_policy_fn(obs, (goal_x, goal_y)) + + return policy_fn diff --git a/wiserl/env/odrl_envs/antmaze/mujoco_goal_env.py b/wiserl/env/odrl_envs/antmaze/mujoco_goal_env.py new file mode 100644 index 0000000..714facb --- /dev/null +++ b/wiserl/env/odrl_envs/antmaze/mujoco_goal_env.py @@ -0,0 +1,191 @@ +from collections import OrderedDict +import os + + +from gym import error, spaces +from gym.utils import seeding +import numpy as np +from os import path +import gym + +try: + import mujoco_py +except ImportError as e: + raise error.DependencyNotInstalled("{}. (HINT: you need to install mujoco_py, and also perform the setup instructions here: https://github.com/openai/mujoco-py/.)".format(e)) + +DEFAULT_SIZE = 500 + +def convert_observation_to_space(observation): + if isinstance(observation, dict): + space = spaces.Dict(OrderedDict([ + (key, convert_observation_to_space(value)) + for key, value in observation.items() + ])) + elif isinstance(observation, np.ndarray): + low = np.full(observation.shape, -float('inf'), dtype=np.float32) + high = np.full(observation.shape, float('inf'), dtype=np.float32) + space = spaces.Box(low, high, dtype=observation.dtype) + else: + raise NotImplementedError(type(observation), observation) + + return space + +class MujocoGoalEnv(gym.Env): + """SuperClass for all MuJoCo goal reaching environments""" + + def __init__(self, model_path, frame_skip): + if model_path.startswith("/"): + fullpath = model_path + else: + fullpath = os.path.join(os.path.dirname(__file__), "assets", model_path) + if not path.exists(fullpath): + raise IOError("File %s does not exist" % fullpath) + self.frame_skip = frame_skip + self.model = mujoco_py.load_model_from_path(fullpath) + self.sim = mujoco_py.MjSim(self.model) + self.data = self.sim.data + self.viewer = None + self._viewers = {} + + self.metadata = { + 'render.modes': ['human', 'rgb_array', 'depth_array'], + 'video.frames_per_second': int(np.round(1.0 / self.dt)) + } + + self.init_qpos = self.sim.data.qpos.ravel().copy() + self.init_qvel = self.sim.data.qvel.ravel().copy() + + self._set_action_space() + + action = self.action_space.sample() + # import ipdb; ipdb.set_trace() + observation, _reward, done, _info = self.step(action) + assert not done + + self._set_observation_space(observation['observation']) + + self.seed() + + def _set_action_space(self): + bounds = self.model.actuator_ctrlrange.copy().astype(np.float32) + low, high = bounds.T + self.action_space = spaces.Box(low=low, high=high, dtype=np.float32) + return self.action_space + + # def _set_observation_space(self, observation): + # self.observation_space = convert_observation_to_space(observation) + # return self.observation_space + + def _set_observation_space(self, observation): + temp_observation_space = convert_observation_to_space(observation) + self.observation_space = spaces.Dict(dict( + observation=temp_observation_space, + desired_goal=spaces.Box(-np.inf, np.inf, shape=(2,), dtype=np.float32), + achieved_goal=spaces.Box(-np.inf, np.inf, shape=(2,), dtype=np.float32), + )) + return self.observation_space + + def seed(self, seed=None): + self.np_random, seed = seeding.np_random(seed) + return [seed] + + # methods to override: + # ---------------------------- + + def reset_model(self): + """ + Reset the robot degrees of freedom (qpos and qvel). + Implement this in each subclass. + """ + raise NotImplementedError + + def viewer_setup(self): + """ + This method is called when the viewer is initialized. + Optionally implement this method, if you need to tinker with camera position + and so forth. + """ + pass + + def reset(self): + self.sim.reset() + ob = self.reset_model() + return ob + + def set_state(self, qpos, qvel): + assert qpos.shape == (self.model.nq,) and qvel.shape == (self.model.nv,) + old_state = self.sim.get_state() + new_state = mujoco_py.MjSimState(old_state.time, qpos, qvel, + old_state.act, old_state.udd_state) + self.sim.set_state(new_state) + self.sim.forward() + + @property + def dt(self): + return self.model.opt.timestep * self.frame_skip + + def do_simulation(self, ctrl, n_frames): + self.sim.data.ctrl[:] = ctrl + for _ in range(n_frames): + self.sim.step() + + def render(self, + mode='human', + width=DEFAULT_SIZE, + height=DEFAULT_SIZE, + camera_id=None, + camera_name=None): + if mode == 'rgb_array': + if camera_id is not None and camera_name is not None: + raise ValueError("Both `camera_id` and `camera_name` cannot be" + " specified at the same time.") + + no_camera_specified = camera_name is None and camera_id is None + if no_camera_specified: + camera_name = 'track' + + if camera_id is None and camera_name in self.model._camera_name2id: + camera_id = self.model.camera_name2id(camera_name) + + self._get_viewer(mode).render(width, height, camera_id=camera_id) + # window size used for old mujoco-py: + data = self._get_viewer(mode).read_pixels(width, height, depth=False) + # original image is upside-down, so flip it + return data[::-1, :, :] + elif mode == 'depth_array': + self._get_viewer(mode).render(width, height) + # window size used for old mujoco-py: + # Extract depth part of the read_pixels() tuple + data = self._get_viewer(mode).read_pixels(width, height, depth=True)[1] + # original image is upside-down, so flip it + return data[::-1, :] + elif mode == 'human': + self._get_viewer(mode).render() + + def close(self): + if self.viewer is not None: + # self.viewer.finish() + self.viewer = None + self._viewers = {} + + def _get_viewer(self, mode): + self.viewer = self._viewers.get(mode) + if self.viewer is None: + if mode == 'human': + self.viewer = mujoco_py.MjViewer(self.sim) + elif mode == 'rgb_array' or mode == 'depth_array': + self.viewer = mujoco_py.MjRenderContextOffscreen(self.sim, -1) + + self.viewer_setup() + self._viewers[mode] = self.viewer + return self.viewer + + def get_body_com(self, body_name): + return self.data.get_body_xpos(body_name) + + def state_vector(self): + return np.concatenate([ + self.sim.data.qpos.flat, + self.sim.data.qvel.flat + ]) + diff --git a/wiserl/env/odrl_envs/antmaze/wrappers.py b/wiserl/env/odrl_envs/antmaze/wrappers.py new file mode 100644 index 0000000..45b371c --- /dev/null +++ b/wiserl/env/odrl_envs/antmaze/wrappers.py @@ -0,0 +1,168 @@ +import numpy as np +import itertools +from gym import Env +from gym.spaces import Box +from gym.spaces import Discrete + +from collections import deque + + +class ProxyEnv(Env): + def __init__(self, wrapped_env): + self._wrapped_env = wrapped_env + self.action_space = self._wrapped_env.action_space + self.observation_space = self._wrapped_env.observation_space + + @property + def wrapped_env(self): + return self._wrapped_env + + def reset(self, **kwargs): + return self._wrapped_env.reset(**kwargs) + + def step(self, action): + return self._wrapped_env.step(action) + + def render(self, *args, **kwargs): + return self._wrapped_env.render(*args, **kwargs) + + @property + def horizon(self): + return self._wrapped_env.horizon + + def terminate(self): + if hasattr(self.wrapped_env, "terminate"): + self.wrapped_env.terminate() + + def __getattr__(self, attr): + if attr == '_wrapped_env': + raise AttributeError() + return getattr(self._wrapped_env, attr) + + def __getstate__(self): + """ + This is useful to override in case the wrapped env has some funky + __getstate__ that doesn't play well with overriding __getattr__. + + The main problematic case is/was gym's EzPickle serialization scheme. + :return: + """ + return self.__dict__ + + def __setstate__(self, state): + self.__dict__.update(state) + + def __str__(self): + return '{}({})'.format(type(self).__name__, self.wrapped_env) + + +class HistoryEnv(ProxyEnv, Env): + def __init__(self, wrapped_env, history_len): + super().__init__(wrapped_env) + self.history_len = history_len + + high = np.inf * np.ones( + self.history_len * self.observation_space.low.size) + low = -high + self.observation_space = Box(low=low, + high=high, + ) + self.history = deque(maxlen=self.history_len) + + def step(self, action): + state, reward, done, info = super().step(action) + self.history.append(state) + flattened_history = self._get_history().flatten() + return flattened_history, reward, done, info + + def reset(self, **kwargs): + state = super().reset() + self.history = deque(maxlen=self.history_len) + self.history.append(state) + flattened_history = self._get_history().flatten() + return flattened_history + + def _get_history(self): + observations = list(self.history) + + obs_count = len(observations) + for _ in range(self.history_len - obs_count): + dummy = np.zeros(self._wrapped_env.observation_space.low.size) + observations.append(dummy) + return np.c_[observations] + + +class DiscretizeEnv(ProxyEnv, Env): + def __init__(self, wrapped_env, num_bins): + super().__init__(wrapped_env) + low = self.wrapped_env.action_space.low + high = self.wrapped_env.action_space.high + action_ranges = [ + np.linspace(low[i], high[i], num_bins) + for i in range(len(low)) + ] + self.idx_to_continuous_action = [ + np.array(x) for x in itertools.product(*action_ranges) + ] + self.action_space = Discrete(len(self.idx_to_continuous_action)) + + def step(self, action): + continuous_action = self.idx_to_continuous_action[action] + return super().step(continuous_action) + + +class NormalizedBoxEnv(ProxyEnv): + """ + Normalize action to in [-1, 1]. + + Optionally normalize observations and scale reward. + """ + + def __init__( + self, + env, + reward_scale=1., + obs_mean=None, + obs_std=None, + ): + ProxyEnv.__init__(self, env) + self._should_normalize = not (obs_mean is None and obs_std is None) + if self._should_normalize: + if obs_mean is None: + obs_mean = np.zeros_like(env.observation_space.low) + else: + obs_mean = np.array(obs_mean) + if obs_std is None: + obs_std = np.ones_like(env.observation_space.low) + else: + obs_std = np.array(obs_std) + self._reward_scale = reward_scale + self._obs_mean = obs_mean + self._obs_std = obs_std + ub = np.ones(self._wrapped_env.action_space.shape) + self.action_space = Box(-1 * ub, ub) + + def estimate_obs_stats(self, obs_batch, override_values=False): + if self._obs_mean is not None and not override_values: + raise Exception("Observation mean and std already set. To " + "override, set override_values to True.") + self._obs_mean = np.mean(obs_batch, axis=0) + self._obs_std = np.std(obs_batch, axis=0) + + def _apply_normalize_obs(self, obs): + return (obs - self._obs_mean) / (self._obs_std + 1e-8) + + def step(self, action): + lb = self._wrapped_env.action_space.low + ub = self._wrapped_env.action_space.high + scaled_action = lb + (action + 1.) * 0.5 * (ub - lb) + scaled_action = np.clip(scaled_action, lb, ub) + + wrapped_step = self._wrapped_env.step(scaled_action) + next_obs, reward, done, info = wrapped_step + if self._should_normalize: + next_obs = self._apply_normalize_obs(next_obs) + return next_obs, reward * self._reward_scale, done, info + + def __str__(self): + return "Normalized: %s" % self._wrapped_env diff --git a/wiserl/env/odrl_envs/infos.py b/wiserl/env/odrl_envs/infos.py new file mode 100644 index 0000000..fa14233 --- /dev/null +++ b/wiserl/env/odrl_envs/infos.py @@ -0,0 +1,272 @@ +# reference scores for all benchmark tasks + +REF_MIN_SCORE = { + 'pen-broken-joint-easy' : -12.172796387517222 , + 'pen-broken-joint-medium' : -12.172796387517222 , + 'pen-broken-joint-hard' : -12.172796387517222 , + 'pen-shrink-finger-easy' : -12.172796387517222 , + 'pen-shrink-finger-medium' : -12.172796387517222 , + 'pen-shrink-finger-hard' : -12.172796387517222 , + 'door-broken-joint-easy' : -52.33817104624433 , + 'door-broken-joint-medium' : -52.33817104624433 , + 'door-broken-joint-hard' : -52.33817104624433 , + 'door-shrink-finger-easy' : -52.33817104624433 , + 'door-shrink-finger-medium' : -52.33817104624433 , + 'door-shrink-finger-hard' : -52.33817104624433 , + 'relocate-broken-joint-easy' : -4.439599892829203 , + 'relocate-broken-joint-medium' : -4.439599892829203 , + 'relocate-broken-joint-hard' : -4.439599892829203 , + 'relocate-shrink-finger-easy' : -4.439599892829203 , + 'relocate-shrink-finger-medium' : -4.439599892829203 , + 'relocate-shrink-finger-hard' : -4.439599892829203 , + 'hammer-broken-joint-easy' : -240.92803745715037 , + 'hammer-broken-joint-medium' : -240.92803745715037 , + 'hammer-broken-joint-hard' : -240.92803745715037 , + 'hammer-shrink-finger-easy' : -240.92803745715037 , + 'hammer-shrink-finger-medium' : -240.92803745715037 , + 'hammer-shrink-finger-hard' : -240.92803745715037 , + 'antmaze-small-empty' : 0.0 , + 'antmaze-small-centerblock' : 0.0 , + 'antmaze-small-lshape' : 0.0 , + 'antmaze-small-zshape' : 0.0 , + 'antmaze-small-reverseu' : 0.0 , + 'antmaze-small-reversel' : 0.0 , + 'antmaze-medium-1' : 0.0 , + 'antmaze-medium-2' : 0.0 , + 'antmaze-medium-3' : 0.0 , + 'antmaze-medium-4' : 0.0 , + 'antmaze-medium-5' : 0.0 , + 'antmaze-medium-6' : 0.0 , + 'antmaze-large-1' : 0.0 , + 'antmaze-large-2' : 0.0 , + 'antmaze-large-3' : 0.0 , + 'antmaze-large-4' : 0.0 , + 'antmaze-large-5' : 0.0 , + 'antmaze-large-6' : 0.0 , + 'halfcheetah-friction-0.1' : -280.178953 , + 'halfcheetah-friction-0.5' : -280.178953 , + 'halfcheetah-friction-1.0' : -280.178953 , + 'halfcheetah-friction-2.0' : -280.178953 , + 'halfcheetah-friction-5.0' : -280.178953 , + 'halfcheetah-gravity-0.1' : -280.178953 , + 'halfcheetah-gravity-0.5' : -280.178953 , + 'halfcheetah-gravity-1.0' : -280.178953 , + 'halfcheetah-gravity-2.0' : -280.178953 , + 'halfcheetah-gravity-5.0' : -280.178953 , + 'halfcheetah-kinematic-footjnt-easy': -280.178953 , + 'halfcheetah-kinematic-footjnt-medium': -280.178953 , + 'halfcheetah-kinematic-footjnt-hard': -280.178953 , + 'halfcheetah-kinematic-thighjnt-easy': -280.178953 , + 'halfcheetah-kinematic-thighjnt-medium': -280.178953 , + 'halfcheetah-kinematic-thighjnt-hard': -280.178953 , + 'halfcheetah-morph-thigh-easy': -280.178953 , + 'halfcheetah-morph-thigh-medium': -280.178953 , + 'halfcheetah-morph-thigh-hard': -280.178953 , + 'halfcheetah-morph-torso-easy': -280.178953 , + 'halfcheetah-morph-torso-medium': -280.178953 , + 'halfcheetah-morph-torso-hard': -280.178953 , + 'hopper-friction-0.1' : -26.3360015397715 , + 'hopper-friction-0.5' : -26.3360015397715 , + 'hopper-friction-1.0' : -26.3360015397715 , + 'hopper-friction-2.0' : -26.3360015397715 , + 'hopper-friction-5.0' : -26.3360015397715 , + 'hopper-gravity-0.1' : -26.3360015397715 , + 'hopper-gravity-0.5' : -26.3360015397715 , + 'hopper-gravity-1.0' : -26.3360015397715 , + 'hopper-gravity-2.0' : -26.3360015397715 , + 'hopper-gravity-5.0' : -26.3360015397715 , + 'hopper-kinematic-footjnt-easy': -26.3360015397715 , + 'hopper-kinematic-footjnt-medium': -26.3360015397715 , + 'hopper-kinematic-footjnt-hard': -26.3360015397715 , + 'hopper-kinematic-legjnt-easy': -26.3360015397715 , + 'hopper-kinematic-legjnt-medium': -26.3360015397715 , + 'hopper-kinematic-legjnt-hard': -26.3360015397715 , + 'hopper-morph-foot-easy': -26.3360015397715 , + 'hopper-morph-foot-medium': -26.3360015397715 , + 'hopper-morph-foot-hard': -26.3360015397715 , + 'hopper-morph-torso-easy': -26.3360015397715 , + 'hopper-morph-torso-medium': -26.3360015397715 , + 'hopper-morph-torso-hard': -26.3360015397715 , + 'walker2d-friction-0.1' : 10.079455055289959 , + 'walker2d-friction-0.5' : 10.079455055289959 , + 'walker2d-friction-1.0' : 10.079455055289959 , + 'walker2d-friction-2.0' : 10.079455055289959 , + 'walker2d-friction-5.0' : 10.079455055289959 , + 'walker2d-gravity-0.1' : 10.079455055289959 , + 'walker2d-gravity-0.5' : 10.079455055289959 , + 'walker2d-gravity-1.0' : 10.079455055289959 , + 'walker2d-gravity-2.0' : 10.079455055289959 , + 'walker2d-gravity-5.0' : 10.079455055289959 , + 'walker2d-kinematic-footjnt-easy': 10.079455055289959 , + 'walker2d-kinematic-footjnt-medium': 10.079455055289959 , + 'walker2d-kinematic-footjnt-hard': 10.079455055289959 , + 'walker2d-kinematic-thighjnt-easy': 10.079455055289959 , + 'walker2d-kinematic-thighjnt-medium': 10.079455055289959 , + 'walker2d-kinematic-thighjnt-hard': 10.079455055289959 , + 'walker2d-morph-leg-easy': 10.079455055289959 , + 'walker2d-morph-leg-medium': 10.079455055289959 , + 'walker2d-morph-leg-hard': 10.079455055289959 , + 'walker2d-morph-torso-easy': 10.079455055289959 , + 'walker2d-morph-torso-medium': 10.079455055289959 , + 'walker2d-morph-torso-hard': 10.079455055289959 , + 'ant-friction-0.1' : -325.6 , + 'ant-friction-0.5' : -325.6 , + 'ant-friction-1.0' : -325.6 , + 'ant-friction-2.0' : -325.6 , + 'ant-friction-5.0' : -325.6 , + 'ant-gravity-0.1' : -325.6 , + 'ant-gravity-0.5' : -325.6 , + 'ant-gravity-1.0' : -325.6 , + 'ant-gravity-2.0' : -325.6 , + 'ant-gravity-5.0' : -325.6 , + 'ant-kinematic-anklejnt-easy': -325.6 , + 'ant-kinematic-anklejnt-medium': -325.6 , + 'ant-kinematic-anklejnt-hard': -325.6 , + 'ant-kinematic-hipjnt-easy': -325.6 , + 'ant-kinematic-hipjnt-medium': -325.6 , + 'ant-kinematic-hipjnt-hard': -325.6 , + 'ant-morph-alllegs-easy': -325.6 , + 'ant-morph-alllegs-medium': -325.6 , + 'ant-morph-alllegs-hard': -325.6 , + 'ant-morph-halflegs-easy': -325.6 , + 'ant-morph-halflegs-medium': -325.6 , + 'ant-morph-halflegs-hard': -325.6 , +} + +REF_MAX_SCORE = { + 'pen-broken-joint-easy' : 6408.3837890625 , + 'pen-broken-joint-medium' : 6408.3837890625 , + 'pen-broken-joint-hard' : 6408.3837890625 , + 'pen-shrink-finger-easy' : 6408.3837890625 , + 'pen-shrink-finger-medium' : 6408.3837890625 , + 'pen-shrink-finger-hard' : 6408.3837890625 , + 'door-broken-joint-easy' : 2880.5693087298737 , + 'door-broken-joint-medium' : 2880.5693087298737 , + 'door-broken-joint-hard' : 2880.5693087298737 , + 'door-shrink-finger-easy' : 2880.5693087298737 , + 'door-shrink-finger-medium' : 2880.5693087298737 , + 'door-shrink-finger-hard' : 2880.5693087298737 , + 'relocate-broken-joint-easy' : 4233.877797728884 , + 'relocate-broken-joint-medium' : 4233.877797728884 , + 'relocate-broken-joint-hard' : 4233.877797728884 , + 'relocate-shrink-finger-easy' : 4233.877797728884 , + 'relocate-shrink-finger-medium' : 4233.877797728884 , + 'relocate-shrink-finger-hard' : 4233.877797728884 , + 'hammer-broken-joint-easy' : 12794.134825156867 , + 'hammer-broken-joint-medium' : 12794.134825156867 , + 'hammer-broken-joint-hard' : 12794.134825156867 , + 'hammer-shrink-finger-easy' : 12794.134825156867 , + 'hammer-shrink-finger-medium' : 12794.134825156867 , + 'hammer-shrink-finger-hard' : 12794.134825156867 , + 'antmaze-small-empty' : 1.0 , + 'antmaze-small-centerblock' : 1.0 , + 'antmaze-small-lshape' : 1.0 , + 'antmaze-small-zshape' : 1.0 , + 'antmaze-small-reverseu' : 1.0 , + 'antmaze-small-reversel' : 1.0 , + 'antmaze-medium-1' : 1.0 , + 'antmaze-medium-2' : 1.0 , + 'antmaze-medium-3' : 1.0 , + 'antmaze-medium-4' : 1.0 , + 'antmaze-medium-5' : 1.0 , + 'antmaze-medium-6' : 1.0 , + 'antmaze-large-1' : 1.0 , + 'antmaze-large-2' : 1.0 , + 'antmaze-large-3' : 1.0 , + 'antmaze-large-4' : 1.0 , + 'antmaze-large-5' : 1.0 , + 'antmaze-large-6' : 1.0 , + 'halfcheetah-friction-0.1' : 41696.546875 , + 'halfcheetah-friction-0.5' : 7357.0712890625 , + 'halfcheetah-friction-1.0' : 11255.9677734375 , + 'halfcheetah-friction-2.0' : 11255.9677734375 , + 'halfcheetah-friction-5.0' : 10199.3271484375 , + 'halfcheetah-gravity-0.1' : 2466.85 , + 'halfcheetah-gravity-0.5' : 9509.15 , + 'halfcheetah-gravity-1.0' : 9509.15 , + 'halfcheetah-gravity-2.0' : 9509.15 , + 'halfcheetah-gravity-5.0' : 3756.24 , + 'halfcheetah-kinematic-footjnt-easy': 12135.0 , + 'halfcheetah-kinematic-footjnt-medium': 12135.0 , + 'halfcheetah-kinematic-footjnt-hard': 12135.0 , + 'halfcheetah-kinematic-thighjnt-easy': 12135.0 , + 'halfcheetah-kinematic-thighjnt-medium': 12135.0 , + 'halfcheetah-kinematic-thighjnt-hard': 12135.0 , + 'halfcheetah-morph-thigh-easy': 12135.0 , + 'halfcheetah-morph-thigh-medium': 12135.0 , + 'halfcheetah-morph-thigh-hard': 12135.0 , + 'halfcheetah-morph-torso-easy': 12135.0 , + 'halfcheetah-morph-torso-medium': 12135.0 , + 'halfcheetah-morph-torso-hard': 12135.0 , + 'hopper-friction-0.1' : 3234.3 , + 'hopper-friction-0.5' : 3234.3 , + 'hopper-friction-1.0' : 3234.3 , + 'hopper-friction-2.0' : 3234.3 , + 'hopper-friction-5.0' : 3234.3 , + 'hopper-gravity-0.1' : 3234.3 , + 'hopper-gravity-0.5' : 3234.3 , + 'hopper-gravity-1.0' : 3234.3 , + 'hopper-gravity-2.0' : 3234.3 , + 'hopper-gravity-5.0' : 3234.3 , + 'hopper-kinematic-footjnt-easy': 3234.3 , + 'hopper-kinematic-footjnt-medium': 3234.3 , + 'hopper-kinematic-footjnt-hard': 3234.3 , + 'hopper-kinematic-legjnt-easy': 3234.3 , + 'hopper-kinematic-legjnt-medium': 3234.3 , + 'hopper-kinematic-legjnt-hard': 3234.3 , + 'hopper-morph-foot-easy': 3234.3 , + 'hopper-morph-foot-medium': 3234.3 , + 'hopper-morph-foot-hard': 3234.3 , + 'hopper-morph-torso-easy': 3234.3 , + 'hopper-morph-torso-medium': 3234.3 , + 'hopper-morph-torso-hard': 3234.3 , + 'walker2d-friction-0.1' : 3360.181 , + 'walker2d-friction-0.5' : 4229.348 , + 'walker2d-friction-1.0' : 5180.044 , + 'walker2d-friction-2.0' : 5180.044 , + 'walker2d-friction-5.0' : 4988.835 , + 'walker2d-gravity-0.1' : 2074.904 , + 'walker2d-gravity-0.5' : 5194.713 , + 'walker2d-gravity-1.0' : 5056.445 , + 'walker2d-gravity-2.0' : 5056.445 , + 'walker2d-gravity-5.0' : 3665.385 , + 'walker2d-kinematic-footjnt-easy': 4592.3 , + 'walker2d-kinematic-footjnt-medium': 4592.3 , + 'walker2d-kinematic-footjnt-hard': 4592.3 , + 'walker2d-kinematic-thighjnt-easy': 4592.3 , + 'walker2d-kinematic-thighjnt-medium': 4592.3 , + 'walker2d-kinematic-thighjnt-hard': 4592.3 , + 'walker2d-morph-leg-easy': 4592.3 , + 'walker2d-morph-leg-medium': 4592.3 , + 'walker2d-morph-leg-hard': 4592.3 , + 'walker2d-morph-torso-easy': 4592.3 , + 'walker2d-morph-torso-medium': 4592.3 , + 'walker2d-morph-torso-hard': 4592.3 , + 'ant-friction-0.1' : 7938.962 , + 'ant-friction-0.5' : 8301.338 , + 'ant-friction-1.0' : 5167.376 , + 'ant-friction-2.0' : 5167.376 , + 'ant-friction-5.0' : 4545.021 , + 'ant-gravity-0.1' : 2782.098 , + 'ant-gravity-0.5' : 4317.065 , + 'ant-gravity-1.0' : 6705.12 , + 'ant-gravity-2.0' : 6705.12 , + 'ant-gravity-5.0' : 6226.89 , + 'ant-kinematic-anklejnt-easy': 5139.832 , + 'ant-kinematic-anklejnt-medium': 5139.832 , + 'ant-kinematic-anklejnt-hard': 5139.832 , + 'ant-kinematic-hipjnt-easy': 5139.832 , + 'ant-kinematic-hipjnt-medium': 5139.832 , + 'ant-kinematic-hipjnt-hard': 5139.832 , + 'ant-morph-alllegs-easy': 5139.832 , + 'ant-morph-alllegs-medium': 5139.832 , + 'ant-morph-alllegs-hard': 5139.832 , + 'ant-morph-halflegs-easy': 5139.832 , + 'ant-morph-halflegs-medium': 5139.832 , + 'ant-morph-halflegs-hard': 5139.832 , +} + +def get_normalized_score(score, env_name): + ref_min_score = REF_MIN_SCORE[env_name] + ref_max_score = REF_MAX_SCORE[env_name] + return (score - ref_min_score) / (ref_max_score - ref_min_score) * 100 \ No newline at end of file diff --git a/wiserl/env/odrl_envs/mujoco/__init__.py b/wiserl/env/odrl_envs/mujoco/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant.xml b/wiserl/env/odrl_envs/mujoco/assets/ant.xml new file mode 100644 index 0000000..ee4d679 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant.xml @@ -0,0 +1,81 @@ + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_friction_0.1.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_friction_0.1.xml new file mode 100644 index 0000000..c0866c9 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_friction_0.1.xml @@ -0,0 +1,81 @@ + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_friction_0.5.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_friction_0.5.xml new file mode 100644 index 0000000..b175773 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_friction_0.5.xml @@ -0,0 +1,81 @@ + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_friction_1.0.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_friction_1.0.xml new file mode 100644 index 0000000..ee4d679 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_friction_1.0.xml @@ -0,0 +1,81 @@ + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_friction_2.0.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_friction_2.0.xml new file mode 100644 index 0000000..21549a9 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_friction_2.0.xml @@ -0,0 +1,81 @@ + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_friction_5.0.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_friction_5.0.xml new file mode 100644 index 0000000..3e9a09e --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_friction_5.0.xml @@ -0,0 +1,81 @@ + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_gravity_0.1.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_gravity_0.1.xml new file mode 100644 index 0000000..967b452 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_gravity_0.1.xml @@ -0,0 +1,82 @@ + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_gravity_0.5.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_gravity_0.5.xml new file mode 100644 index 0000000..de9edd3 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_gravity_0.5.xml @@ -0,0 +1,82 @@ + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_gravity_1.0.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_gravity_1.0.xml new file mode 100644 index 0000000..bac0d1a --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_gravity_1.0.xml @@ -0,0 +1,82 @@ + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_gravity_2.0.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_gravity_2.0.xml new file mode 100644 index 0000000..8195d61 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_gravity_2.0.xml @@ -0,0 +1,82 @@ + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_gravity_5.0.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_gravity_5.0.xml new file mode 100644 index 0000000..49950e1 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_gravity_5.0.xml @@ -0,0 +1,82 @@ + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_anklejnt_easy.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_anklejnt_easy.xml new file mode 100644 index 0000000..0772c6b --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_anklejnt_easy.xml @@ -0,0 +1,82 @@ + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_anklejnt_hard.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_anklejnt_hard.xml new file mode 100644 index 0000000..ae7c015 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_anklejnt_hard.xml @@ -0,0 +1,82 @@ + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_anklejnt_medium.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_anklejnt_medium.xml new file mode 100644 index 0000000..385656e --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_anklejnt_medium.xml @@ -0,0 +1,82 @@ + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_hipjnt_easy.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_hipjnt_easy.xml new file mode 100644 index 0000000..51031f4 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_hipjnt_easy.xml @@ -0,0 +1,82 @@ + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_hipjnt_hard.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_hipjnt_hard.xml new file mode 100644 index 0000000..ae83a4a --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_hipjnt_hard.xml @@ -0,0 +1,82 @@ + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_hipjnt_medium.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_hipjnt_medium.xml new file mode 100644 index 0000000..98081a4 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_kinematic_hipjnt_medium.xml @@ -0,0 +1,82 @@ + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_morph_alllegs_easy.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_morph_alllegs_easy.xml new file mode 100644 index 0000000..6b96024 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_morph_alllegs_easy.xml @@ -0,0 +1,82 @@ + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_morph_alllegs_hard.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_morph_alllegs_hard.xml new file mode 100644 index 0000000..b1d2b34 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_morph_alllegs_hard.xml @@ -0,0 +1,82 @@ + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_morph_alllegs_medium.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_morph_alllegs_medium.xml new file mode 100644 index 0000000..cb763e6 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_morph_alllegs_medium.xml @@ -0,0 +1,82 @@ + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_morph_halflegs_easy.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_morph_halflegs_easy.xml new file mode 100644 index 0000000..8055c05 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_morph_halflegs_easy.xml @@ -0,0 +1,82 @@ + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_morph_halflegs_hard.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_morph_halflegs_hard.xml new file mode 100644 index 0000000..628b262 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_morph_halflegs_hard.xml @@ -0,0 +1,82 @@ + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/ant_morph_halflegs_medium.xml b/wiserl/env/odrl_envs/mujoco/assets/ant_morph_halflegs_medium.xml new file mode 100644 index 0000000..2ee0435 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/ant_morph_halflegs_medium.xml @@ -0,0 +1,82 @@ + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/half_cheetah.xml b/wiserl/env/odrl_envs/mujoco/assets/half_cheetah.xml new file mode 100644 index 0000000..338c2e8 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/half_cheetah.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_friction_0.1.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_friction_0.1.xml new file mode 100644 index 0000000..d67da4d --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_friction_0.1.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_friction_0.5.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_friction_0.5.xml new file mode 100644 index 0000000..4675e67 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_friction_0.5.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_friction_1.0.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_friction_1.0.xml new file mode 100644 index 0000000..338c2e8 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_friction_1.0.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_friction_2.0.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_friction_2.0.xml new file mode 100644 index 0000000..39c70c9 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_friction_2.0.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_friction_5.0.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_friction_5.0.xml new file mode 100644 index 0000000..c7b2ddd --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_friction_5.0.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_gravity_0.1.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_gravity_0.1.xml new file mode 100644 index 0000000..b89c564 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_gravity_0.1.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_gravity_0.5.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_gravity_0.5.xml new file mode 100644 index 0000000..97e5276 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_gravity_0.5.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_gravity_1.0.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_gravity_1.0.xml new file mode 100644 index 0000000..338c2e8 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_gravity_1.0.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_gravity_2.0.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_gravity_2.0.xml new file mode 100644 index 0000000..e0d90ba --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_gravity_2.0.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_gravity_5.0.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_gravity_5.0.xml new file mode 100644 index 0000000..85696b4 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_gravity_5.0.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_footjnt_easy.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_footjnt_easy.xml new file mode 100644 index 0000000..e616e27 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_footjnt_easy.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_footjnt_hard.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_footjnt_hard.xml new file mode 100644 index 0000000..af21975 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_footjnt_hard.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_footjnt_medium.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_footjnt_medium.xml new file mode 100644 index 0000000..b859d68 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_footjnt_medium.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_thighjnt_easy.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_thighjnt_easy.xml new file mode 100644 index 0000000..98d4b64 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_thighjnt_easy.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_thighjnt_hard.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_thighjnt_hard.xml new file mode 100644 index 0000000..fc8680e --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_thighjnt_hard.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_thighjnt_medium.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_thighjnt_medium.xml new file mode 100644 index 0000000..545d6fc --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_kinematic_thighjnt_medium.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_thigh_easy.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_thigh_easy.xml new file mode 100644 index 0000000..a0fa74a --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_thigh_easy.xml @@ -0,0 +1,104 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_thigh_hard.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_thigh_hard.xml new file mode 100644 index 0000000..d9af5c3 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_thigh_hard.xml @@ -0,0 +1,104 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_thigh_medium.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_thigh_medium.xml new file mode 100644 index 0000000..29db2f2 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_thigh_medium.xml @@ -0,0 +1,104 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_torso_easy.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_torso_easy.xml new file mode 100644 index 0000000..c55e170 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_torso_easy.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_torso_hard.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_torso_hard.xml new file mode 100644 index 0000000..6210cc7 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_torso_hard.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_torso_medium.xml b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_torso_medium.xml new file mode 100644 index 0000000..37284e5 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/halfcheetah_morph_torso_medium.xml @@ -0,0 +1,96 @@ + + + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper.xml new file mode 100644 index 0000000..f18bc46 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_friction_0.1.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_friction_0.1.xml new file mode 100644 index 0000000..297306a --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_friction_0.1.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_friction_0.5.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_friction_0.5.xml new file mode 100644 index 0000000..4e4090d --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_friction_0.5.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_friction_1.0.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_friction_1.0.xml new file mode 100644 index 0000000..f18bc46 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_friction_1.0.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_friction_2.0.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_friction_2.0.xml new file mode 100644 index 0000000..7979d52 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_friction_2.0.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_friction_5.0.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_friction_5.0.xml new file mode 100644 index 0000000..445368d --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_friction_5.0.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_gravity_0.1.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_gravity_0.1.xml new file mode 100644 index 0000000..54c743b --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_gravity_0.1.xml @@ -0,0 +1,49 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_gravity_0.5.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_gravity_0.5.xml new file mode 100644 index 0000000..43d29c4 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_gravity_0.5.xml @@ -0,0 +1,49 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_gravity_1.0.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_gravity_1.0.xml new file mode 100644 index 0000000..0350976 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_gravity_1.0.xml @@ -0,0 +1,49 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_gravity_2.0.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_gravity_2.0.xml new file mode 100644 index 0000000..2f5ff03 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_gravity_2.0.xml @@ -0,0 +1,49 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_gravity_5.0.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_gravity_5.0.xml new file mode 100644 index 0000000..812f730 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_gravity_5.0.xml @@ -0,0 +1,49 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_footjnt_easy.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_footjnt_easy.xml new file mode 100644 index 0000000..226e154 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_footjnt_easy.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_footjnt_hard.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_footjnt_hard.xml new file mode 100644 index 0000000..3b34007 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_footjnt_hard.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_footjnt_medium.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_footjnt_medium.xml new file mode 100644 index 0000000..ecbb649 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_footjnt_medium.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_legjnt_easy.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_legjnt_easy.xml new file mode 100644 index 0000000..f6f3804 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_legjnt_easy.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_legjnt_hard.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_legjnt_hard.xml new file mode 100644 index 0000000..ad3d1c4 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_legjnt_hard.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_legjnt_medium.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_legjnt_medium.xml new file mode 100644 index 0000000..12fe332 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_kinematic_legjnt_medium.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_foot_easy.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_foot_easy.xml new file mode 100644 index 0000000..28e8936 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_foot_easy.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_foot_hard.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_foot_hard.xml new file mode 100644 index 0000000..d40b2d0 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_foot_hard.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_foot_medium.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_foot_medium.xml new file mode 100644 index 0000000..8d104ec --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_foot_medium.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_torso_easy.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_torso_easy.xml new file mode 100644 index 0000000..17c2fa1 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_torso_easy.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_torso_hard.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_torso_hard.xml new file mode 100644 index 0000000..d26f707 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_torso_hard.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_torso_medium.xml b/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_torso_medium.xml new file mode 100644 index 0000000..bbe7ba6 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/hopper_morph_torso_medium.xml @@ -0,0 +1,48 @@ + + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d.xml new file mode 100644 index 0000000..3342571 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_friction_0.1.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_friction_0.1.xml new file mode 100644 index 0000000..8f3abb5 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_friction_0.1.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_friction_0.5.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_friction_0.5.xml new file mode 100644 index 0000000..1415303 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_friction_0.5.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_friction_1.0.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_friction_1.0.xml new file mode 100644 index 0000000..01cc405 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_friction_1.0.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_friction_2.0.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_friction_2.0.xml new file mode 100644 index 0000000..9215cf6 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_friction_2.0.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_friction_5.0.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_friction_5.0.xml new file mode 100644 index 0000000..6227b30 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_friction_5.0.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_gravity_0.1.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_gravity_0.1.xml new file mode 100644 index 0000000..8ce643a --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_gravity_0.1.xml @@ -0,0 +1,63 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_gravity_0.5.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_gravity_0.5.xml new file mode 100644 index 0000000..ed47389 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_gravity_0.5.xml @@ -0,0 +1,63 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_gravity_1.0.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_gravity_1.0.xml new file mode 100644 index 0000000..b9b85ea --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_gravity_1.0.xml @@ -0,0 +1,63 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_gravity_2.0.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_gravity_2.0.xml new file mode 100644 index 0000000..d8f01ff --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_gravity_2.0.xml @@ -0,0 +1,63 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_gravity_5.0.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_gravity_5.0.xml new file mode 100644 index 0000000..06cd095 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_gravity_5.0.xml @@ -0,0 +1,63 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_footjnt_easy.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_footjnt_easy.xml new file mode 100644 index 0000000..17d3b80 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_footjnt_easy.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_footjnt_hard.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_footjnt_hard.xml new file mode 100644 index 0000000..edbeaf8 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_footjnt_hard.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_footjnt_medium.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_footjnt_medium.xml new file mode 100644 index 0000000..6ece30e --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_footjnt_medium.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_thighjnt_easy.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_thighjnt_easy.xml new file mode 100644 index 0000000..4516e92 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_thighjnt_easy.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_thighjnt_hard.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_thighjnt_hard.xml new file mode 100644 index 0000000..fd9c16c --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_thighjnt_hard.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_thighjnt_medium.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_thighjnt_medium.xml new file mode 100644 index 0000000..859512e --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_kinematic_thighjnt_medium.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_leg_easy.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_leg_easy.xml new file mode 100644 index 0000000..9030403 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_leg_easy.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_leg_hard.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_leg_hard.xml new file mode 100644 index 0000000..bd04fb4 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_leg_hard.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_leg_medium.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_leg_medium.xml new file mode 100644 index 0000000..1411724 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_leg_medium.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_torso_easy.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_torso_easy.xml new file mode 100644 index 0000000..641b5b3 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_torso_easy.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_torso_hard.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_torso_hard.xml new file mode 100644 index 0000000..00e0e58 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_torso_hard.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_torso_medium.xml b/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_torso_medium.xml new file mode 100644 index 0000000..1f4de28 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/assets/walker2d_morph_torso_medium.xml @@ -0,0 +1,62 @@ + + + + + + + diff --git a/wiserl/env/odrl_envs/mujoco/call_mujoco_env.py b/wiserl/env/odrl_envs/mujoco/call_mujoco_env.py new file mode 100644 index 0000000..9d147e0 --- /dev/null +++ b/wiserl/env/odrl_envs/mujoco/call_mujoco_env.py @@ -0,0 +1,113 @@ +from typing import Dict +from pathlib import Path +import gym + +from gym.envs.mujoco.half_cheetah_v3 import HalfCheetahEnv +from gym.envs.mujoco.ant_v3 import AntEnv +from gym.envs.mujoco.walker2d_v3 import Walker2dEnv +from gym.envs.mujoco.hopper_v3 import HopperEnv + +from gym.wrappers.time_limit import TimeLimit + + +def call_mujoco_env(env_name: str) -> gym.Env: + # env_name = env_config['env_name'].lower() # eg. "hopper_friction" or "hopper_morph_foot", body_pard is required only in "morph" or 'kinematic" mode + # shift_level = env_config['shift_level'] # either float(0.1/0.5/...) or level(easy/medium/hard) + shift_level = env_name.split('-')[-1] + env_name = env_name[:-len(shift_level)-1] + + if '-' in env_name: + env_name = env_name.replace('-', '_') + + # assert the shift level legal + if 'morph' in env_name or 'kinematic' in env_name: + assert shift_level in ['easy', 'medium', 'hard'], 'The required shift is not available yet, please consider modify the xml file on your own or use the shift scale among easy, medium, hard' + if 'friction' in env_name or 'gravity' in env_name: + assert float(shift_level) in [0.1, 0.5, 1.0, 2.0, 5.0], 'The required shift is not available yet, please consider modify the xml file on your own or use the shift scale among 0.1, 0.5, 1.0, 2.0, 5.0' + + # decide which task it is, support the following tasks + # hopper/half_cheetah/walker2d/ant - friction + # - gravity + # - morph + # - noise + # - broken + + if 'hopper' in env_name: + if env_name == 'hopper': + return gym.make('Hopper-v3') + elif 'friction' in env_name or 'gravity' in env_name: + return TimeLimit( + HopperEnv(xml_file=f"{str(Path(__file__).parent.absolute())}/assets/{env_name}_{float(shift_level)}.xml",), + max_episode_steps=1000 + ) + elif 'noise' in env_name: + # todo: the modification is directly applied on the executed action, no need to modify the xml file itself + return gym.make('Hopper-v3') + elif 'morph' in env_name or 'kinematic' in env_name: + return TimeLimit( + HopperEnv(xml_file=f"{str(Path(__file__).parent.absolute())}/assets/{env_name}_{shift_level}.xml",), + max_episode_steps=1000 + ) + else: + # print("env_name {env_name} is illegal or not implemented") + raise NotImplementedError + elif "halfcheetah" in env_name: + if env_name == 'halfcheetah': + return gym.make('HalfCheetah-v3') + elif 'friction' in env_name or 'gravity' in env_name: + return TimeLimit( + HalfCheetahEnv(xml_file=f"{str(Path(__file__).parent.absolute())}/assets/{env_name}_{float(shift_level)}.xml",), + max_episode_steps=1000 + ) + elif 'noise' in env_name: + # todo: the modification is directly applied on the executed action, no need to modify the xml file itself + return gym.make('HalfCheetah-v3') + elif 'morph' in env_name or 'kinematic' in env_name: + return TimeLimit( + HalfCheetahEnv(xml_file=f"{str(Path(__file__).parent.absolute())}/assets/{env_name}_{shift_level}.xml",), + max_episode_steps=1000 + ) + else: + # print("env_name {env_name} is illegal or not implemented") + raise NotImplementedError + elif "walker2d" in env_name: + if env_name == 'walker2d': + return gym.make('Walker2d-v3') + elif 'friction' in env_name or 'gravity' in env_name: + return TimeLimit( + Walker2dEnv(xml_file=f"{str(Path(__file__).parent.absolute())}/assets/{env_name}_{float(shift_level)}.xml",), + max_episode_steps=1000 + ) + elif 'noise' in env_name: + # todo: the modification is directly applied on the executed action, no need to modify the xml file itself + return gym.make('Walker2d-v3') + elif 'morph' in env_name or 'kinematic' in env_name: + return TimeLimit( + Walker2dEnv(xml_file=f"{str(Path(__file__).parent.absolute())}/assets/{env_name}_{shift_level}.xml",), + max_episode_steps=1000 + ) + else: + # print("env_name {env_name} is illegal or not implemented") + raise NotImplementedError + elif 'ant' in env_name: + if env_name == 'ant': + return gym.make('Ant-v3') + elif 'friction' in env_name or 'gravity' in env_name: + return TimeLimit( + AntEnv(xml_file=f"{str(Path(__file__).parent.absolute())}/assets/{env_name}_{float(shift_level)}.xml",), + max_episode_steps=1000 + ) + elif 'noise' in env_name: + # todo: the modification is directly applied on the executed action, no need to modify the xml file itself + return gym.make('Ant-v3') + elif 'morph' in env_name or 'kinematic' in env_name: + return TimeLimit( + AntEnv(xml_file=f"{str(Path(__file__).parent.absolute())}/assets/{env_name}_{shift_level}.xml",), + max_episode_steps=1000 + ) + else: + # print("env_name {env_name} is illegal or not implemented") + raise NotImplementedError + else: + # print("env_name {env_name} is illegal or not implemented") + raise NotImplementedError \ No newline at end of file diff --git a/wiserl/eval/__init__.py b/wiserl/eval/__init__.py index f422d06..d5d2b60 100644 --- a/wiserl/eval/__init__.py +++ b/wiserl/eval/__init__.py @@ -2,6 +2,5 @@ from wiserl.eval.offline import eval_offline from wiserl.eval.reward_model import eval_reward_model - def eval_placeholder(*args, **kwargs): return {} diff --git a/wiserl/eval/reward_model.py b/wiserl/eval/reward_model.py index 10c6e22..cb0ac47 100644 --- a/wiserl/eval/reward_model.py +++ b/wiserl/eval/reward_model.py @@ -8,7 +8,7 @@ import wiserl.dataset from wiserl.algorithm.base import Algorithm - +from scipy import stats @torch.no_grad() def eval_reward_model( @@ -25,19 +25,30 @@ def eval_reward_model( env.action_space, **kwargs ) + rewards = [] + true_rewards = [] for batch in eval_dataset.create_sequential_iter(): batch = algorithm.format_batch(batch) batch["obs"] = torch.concat([batch["obs_1"], batch["obs_2"]], dim=0) batch["action"] = torch.concat([batch["action_1"], batch["action_2"]], dim=0) reward = algorithm.select_reward(batch) r1, r2 = torch.chunk(reward, 2, dim=0) + + rewards.extend(r1.sum(dim=1).squeeze().cpu()) + rewards.extend(r2.sum(dim=1).squeeze().cpu()) + true_rewards.extend(batch['reward_1'].sum(dim=1).cpu()) + true_rewards.extend(batch['reward_2'].sum(dim=1).cpu()) + logit = r2.sum(dim=1) - r1.sum(dim=1) label = batch["label"].float() reward_loss = algorithm.reward_criterion(logit, label) reward_acc = ((logit > 0) == torch.round(label)).float() rm_eval_loss.extend(reward_loss) rm_eval_acc.extend(reward_acc) + + r, _ = stats.pearsonr(rewards, true_rewards) return { "val_loss": torch.as_tensor(rm_eval_loss).mean().item(), "val_acc": torch.as_tensor(rm_eval_acc).mean().item(), + "val_pearsonr": r } diff --git a/wiserl/trainer/offline_trainer.py b/wiserl/trainer/offline_trainer.py index f36b3c5..5afaed8 100644 --- a/wiserl/trainer/offline_trainer.py +++ b/wiserl/trainer/offline_trainer.py @@ -30,6 +30,7 @@ def __init__( eval_freq: int = 1000, profile_freq: int = -1, checkpoint_freq: Optional[int] = None, + normalize_reward: bool = False, logger: Optional[BaseLogger] = None, device: Union[str, torch.device] = "cpu" ): @@ -49,6 +50,7 @@ def __init__( self.eval_freq = eval_freq self.profile_freq = profile_freq self.checkpoint_freq = checkpoint_freq + self.normalize_reward = normalize_reward self.logger = logger # Datasets and dataloaders @@ -79,6 +81,9 @@ def eval_env(self): def train(self): self.logger.info("Set up datasets and dataloaders") self._datasets = self.setup_datasets(self.dataset_kwargs) + if self.normalize_reward: + for d in self._datasets: + d.normalize_reward() self._dataloaders, self._dataloaders_iter = self.setup_dataloaders(self._datasets, self.dataloader_kwargs) # start training diff --git a/wiserl/trainer/online_trainer.py b/wiserl/trainer/online_trainer.py new file mode 100644 index 0000000..f020953 --- /dev/null +++ b/wiserl/trainer/online_trainer.py @@ -0,0 +1,277 @@ +import os +import random +import tempfile +import time +from collections import defaultdict +from typing import Any, Callable, Dict, List, Optional, Sequence, Union + +import gym +import numpy as np +import torch +from tqdm import trange +from UtilsRL.logger import BaseLogger +from UtilsRL.rl.buffer import TransitionSimpleReplay + +import wiserl.dataset +import wiserl.eval + + +class OnlineTrainer(object): + def __init__( + self, + algorithm, + env_fn: Optional[Callable] = None, + eval_env_fn: Optional[Callable] = None, + eval_kwargs: Optional[dict] = None, + max_buffer_size: int = 100000, + batch_size: int = 256, + buffer_kwargs: Optional[Sequence[Dict]] = None, + random_policy_step: int = 5000, + warmup_step: int = 2000, + max_trajectory_length: int = 1000, + total_steps: int = 1000, + log_freq: int = 100, + env_freq: int = 1, + eval_freq: int = 1000, + profile_freq: int = -1, + checkpoint_freq: Optional[int] = None, + logger: Optional[BaseLogger] = None, + device: Union[str, torch.device] = "cpu" + ): + # The base model + self.algorithm = algorithm + + # Environment parameters + self._env = None + self.env_fn = env_fn + self._eval_env = None + self.eval_env_fn = eval_env_fn + + # Logging parameters + self.warmup_step = warmup_step + self.random_policy_step = random_policy_step + self.max_trajectory_length =max_trajectory_length + self.total_steps = total_steps + self.log_freq = log_freq + self.env_freq = env_freq + self.eval_freq = eval_freq + self.profile_freq = profile_freq + self.checkpoint_freq = checkpoint_freq + self.logger = logger + + # Buffer + self.max_buffer_size = max_buffer_size + self.batch_size = batch_size + self.buffer_kwargs = buffer_kwargs + + # evaluation + self.eval_fn = None + self.eval_kwargs = eval_kwargs + + self.algorithm = self.algorithm.to(device) + self.device = device + + @property + def env(self): + if self._env is None and self.env_fn is not None: + self._env = self.env_fn() + return self._env + + @property + def eval_env(self): + if self._eval_env is None and self.eval_env_fn is not None: + self._eval_env = self.eval_env_fn() + return self._eval_env + + def train(self): + self.logger.info("Set up buffer") + self.buffer = self.setup_buffer(self.max_buffer_size, self.buffer_kwargs) + + # warm up + self.logger.info("Warm up") + obs, terminal = self.env.reset(), False + cur_traj_length = 0 + for step in trange(0, self.warmup_step+1): + action = self.env.action_space.sample() + next_obs, reward, terminal, info = self.env.step(action) + cur_traj_length += 1 + if cur_traj_length >= self.max_trajectory_length: + terminal = False + self.buffer.add_sample( + { + "obs": obs, + "action": action, + "next_obs": next_obs, + "reward": reward, + "terminal": terminal, + } + ) + obs = next_obs + if terminal or cur_traj_length >= self.max_trajectory_length: + obs = self.env.reset() + cur_traj_length = 0 + + # start training + self.logger.info("Start Training") + self.algorithm.train() + + if self.env_freq is not None: + env_freq = int(self.env_freq) if self.env_freq >= 1 else 1 + env_iters = int(1 / self.env_freq) if self.env_freq < 1 else 1 + else: + env_freq = env_iters = None + + obs, terminal = self.env.reset(), False + cur_traj_length = 0 + for step in trange(0, self.total_steps+1): + # do env step + if env_freq and step % env_freq == 0: + for _ in range(env_iters): + self.algorithm.env_step(self.env, step, self.total_steps) + + if step < self.random_policy_step + 1: + action = self.env.action_space.sample() + else: + batch = dict(obs=obs) + batch = self.algorithm.format_batch(batch) + with torch.no_grad(): + action = self.algorithm.select_action(batch) + + next_obs, reward, terminal, info = self.env.step(action) + cur_traj_length += 1 + if cur_traj_length >= self.max_trajectory_length: + terminal = False + self.buffer.add_sample( + { + "obs": obs, + "action": action, + "next_obs": next_obs, + "reward": reward, + "terminal": terminal, + } + ) + obs = next_obs + if terminal or cur_traj_length >= self.max_trajectory_length: + obs = self.env.reset() + cur_traj_length = 0 + + # do algorithm train step + batches = self.buffer.random_batch(self.batch_size) + batches = self.algorithm.format_batch(batches) + metrics = self.algorithm.train_step(batches, step=step, total_steps=self.total_steps) + + # log the metrics + if step % self.log_freq == 0: + self.logger.log_scalars("", metrics, step=step) + + # run eval and validation + if self.eval_freq and step % self.eval_freq == 0: + self.algorithm.eval() + eval_metrics = self.evaluate() + self.logger.log_scalars("eval", eval_metrics, step=step) + self.algorithm.train() + + if self.checkpoint_freq and step % self.checkpoint_freq == 0: + checkpoint_metadata = dict(step=step) + self.algorithm.save(self.logger.output_dir, f"step_{step}.pt", checkpoint_metadata) + + # clean up + checkpoint_metadata = dict(step=step) + self.algorithm.save(self.logger.output_dir, "final.pt", checkpoint_metadata) + if self._env is not None: + self._env.close() + if self._eval_env is not None: + self._eval_env.close() + + def setup_datasets(self, dataset_kwargs=None): + observation_space = self.env.observation_space + action_space = self.env.action_space + + # parse the dataset arguments + if dataset_kwargs is None: + return + if isinstance(dataset_kwargs, dict): + dataset_kwargs = [dataset_kwargs, ] + elif not isinstance(dataset_kwargs, list): + raise TypeError(f"The type of dataset_kwargs should be either list or dict.") + + _datasets = [] + for kwargs in dataset_kwargs: + cls = kwargs.pop("class") + ds = vars(wiserl.dataset)[cls]( + observation_space, + action_space, + **kwargs + ) + _datasets.append(ds) + + return _datasets + + def setup_dataloaders(self, datasets, dataloader_kwargs=None): + if dataloader_kwargs is None: + dataloader_kwargs = {} + if isinstance(dataloader_kwargs, dict): + dataloader_kwargs = [dataloader_kwargs.copy() for _ in range(len(datasets))] + elif not isinstance(dataloader_kwargs, list): + raise TypeError(f"The type of dataloader kwargs should be either dict or list.") + + _dataloaders = [] + for ds, dl_kwargs in zip(datasets, dataloader_kwargs): + _dataloaders.append(torch.utils.data.DataLoader(ds, **dl_kwargs)) + + _dataloaders_iter = [iter(dl) for dl in _dataloaders] + return _dataloaders, _dataloaders_iter + + def setup_buffer(self, max_buffer_size, buffer_kwargs=None): + obs_shape = self.env.observation_space.shape[0] + action_shape = self.env.action_space.shape[-1] + buffer = TransitionSimpleReplay( + max_size=max_buffer_size, + field_specs={ + "obs": { + "shape": [ + obs_shape, + ], + "dtype": np.float32, + }, + "action": { + "shape": [ + action_shape, + ], + "dtype": np.float32, + }, + "next_obs": { + "shape": [ + obs_shape, + ], + "dtype": np.float32, + }, + "reward": { + "shape": [ + 1, + ], + "dtype": np.float32, + }, + "terminal": { + "shape": [ + 1, + ], + "dtype": np.float32, + }, + }, + #**buffer_kwargs + ) + buffer.reset() + return buffer + + def evaluate(self): + assert not self.algorithm.training + if self.eval_kwargs is None: + return {} + if self.eval_fn is None: + self.eval_fn = vars(wiserl.eval)[self.eval_kwargs.pop("function")] + eval_metrics = self.eval_fn( + self.eval_env, self.algorithm, + **self.eval_kwargs + ) + return eval_metrics diff --git a/wiserl/trainer/rmb_offline_trainer.py b/wiserl/trainer/rmb_offline_trainer.py index 91525cb..a03f31c 100644 --- a/wiserl/trainer/rmb_offline_trainer.py +++ b/wiserl/trainer/rmb_offline_trainer.py @@ -30,7 +30,6 @@ def __init__( rl_dataloader_kwargs: Optional[Sequence[Dict]] = None, rl_steps: int = 1000, rl_eval_kwargs: Optional[dict] = None, - rm_label: bool=False, load_rm_path: Optional[str] = None, save_rm_path: Optional[str] = None, log_freq: int = 100, @@ -38,6 +37,8 @@ def __init__( eval_freq: int = 1000, profile_freq: int = -1, checkpoint_freq: Optional[int] = None, + label_reward: bool = True, + normalize_reward: bool = False, logger: Optional[BaseLogger] = None, device: Union[str, torch.device] = "cpu" ): @@ -54,12 +55,14 @@ def __init__( eval_freq=eval_freq, profile_freq=profile_freq, checkpoint_freq=checkpoint_freq, + normalize_reward=normalize_reward, logger=logger, device=device ) self.rm_steps = rm_steps self.rl_steps = rl_steps - self.rm_label = rm_label + self.label_reward = label_reward + self.normalize_reward = normalize_reward self.load_rm_path = load_rm_path self.save_rm_path = save_rm_path # rm & rl datasets, dataloaders, and evals @@ -108,12 +111,15 @@ def train(self): # finally train the rl agent self.logger.info(f"Setting up rl datasets and dataloaders ...") self._rl_datasets = self.setup_datasets(self.rl_dataset_kwargs) - if self.rm_label: + if self.label_reward: self.logger.info(f"Relabeling the reward using pretrained reward model ...") self.algorithm.eval() for d in self._rl_datasets: d.relabel_reward(self.algorithm) self.algorithm.train() + if self.normalize_reward: + for d in self._rl_datasets: + d.normalize_reward() self._rl_dataloaders, self._rl_dataloaders_iter = self.setup_dataloaders(self._rl_datasets, self.rl_dataloader_kwargs) for step in trange(0, self.rl_steps+1, desc="RL"): batches = [next(d) for d in self._rl_dataloaders_iter] diff --git a/wiserl/trainer/rmb_online_trainer.py b/wiserl/trainer/rmb_online_trainer.py new file mode 100644 index 0000000..f9a6723 --- /dev/null +++ b/wiserl/trainer/rmb_online_trainer.py @@ -0,0 +1,249 @@ +import os +import random +import tempfile +import time +from collections import defaultdict +from typing import Any, Callable, Dict, List, Optional, Sequence, Union + +import gym +import numpy as np +import torch +from tqdm import trange +from UtilsRL.logger import BaseLogger + +import wiserl.dataset +import wiserl.eval +from wiserl.trainer.online_trainer import OnlineTrainer + + +class RewardModelBasedOnlineTrainer(OnlineTrainer): + def __init__( + self, + algorithm, + env_fn: Optional[Callable] = None, + eval_env_fn: Optional[Callable] = None, + rm_dataset_kwargs: Optional[Sequence[str]] = None, + rm_dataloader_kwargs: Optional[Sequence[Dict]] = None, + rm_steps: int = 1000, + rm_eval_kwargs: Optional[dict] = None, + max_buffer_size: int = 100000, + batch_size: int = 256, + buffer_kwargs: Optional[Sequence[Dict]] = None, + random_policy_step: int = 5000, + warmup_step: int = 2000, + max_trajectory_length: int = 1000, + rl_steps: int = 1000, + rl_eval_kwargs: Optional[dict] = None, + rm_label: bool=False, + load_rm_path: Optional[str] = None, + save_rm_path: Optional[str] = None, + log_freq: int = 100, + env_freq: int = 1, + eval_freq: int = 1000, + profile_freq: int = -1, + checkpoint_freq: Optional[int] = None, + logger: Optional[BaseLogger] = None, + device: Union[str, torch.device] = "cpu" + ): + super().__init__( + algorithm=algorithm, + env_fn=env_fn, + eval_env_fn=eval_env_fn, + eval_kwargs=None, + max_buffer_size=max_buffer_size, + batch_size=batch_size, + buffer_kwargs=buffer_kwargs, + random_policy_step=random_policy_step, + warmup_step=warmup_step, + max_trajectory_length=max_trajectory_length, + total_steps=rl_steps, + log_freq=log_freq, + env_freq=env_freq, + eval_freq=eval_freq, + profile_freq=profile_freq, + checkpoint_freq=checkpoint_freq, + logger=logger, + device=device + ) + self.rm_steps = rm_steps + self.rl_steps = rl_steps + self.rm_label = rm_label + self.load_rm_path = load_rm_path + self.save_rm_path = save_rm_path + # rm & rl datasets, dataloaders, and evals + self._rm_datasets = self._rl_datasets = None + self._rm_dataloaders = self._rl_dataloaders = None + self._rm_eval_fn = self._rl_eval_fm = None + + self.rm_dataset_kwargs = rm_dataset_kwargs + self.rm_dataloader_kwargs = rm_dataloader_kwargs + self.rm_eval_kwargs = rm_eval_kwargs + self.rl_eval_kwargs = rl_eval_kwargs + + def train(self): + self.algorithm.train() + + # first train the reward model + if self.load_rm_path is not None: + self.logger.info(f"Loading pretrained model from {self.load_rm_path} ... ") + self.algorithm.load_pretrain(self.load_rm_path) + else: + self.logger.info("Setting up pretrain datasets and dataloaders ... ") + self._rm_datasets = self.setup_datasets(self.rm_dataset_kwargs) + self._rm_dataloaders, self._rm_dataloaders_iter = self.setup_dataloaders(self._rm_datasets, self.rm_dataloader_kwargs) + + self.logger.info("Starting pretraining ... ") + for step in trange(0, self.rm_steps+1, desc="pretrain"): + batches = [next(d) for d in self._rm_dataloaders_iter] + batches = self.algorithm.format_batch(batches) + pretrain_metrics = self.algorithm.pretrain_step(batches, step=step, total_steps=self.rm_steps) + + if step % self.log_freq == 0: + self.logger.log_scalars("pretrain", pretrain_metrics, step=step) + + if self.eval_freq and step % self.eval_freq == 0: + self.algorithm.eval() + eval_metrics = self.rm_evaluate() + self.logger.log_scalars("eval", eval_metrics, step=step) + self.algorithm.train() + + if self.save_rm_path is not None: + self.logger.info(f"Saving pretrained model to {self.save_rm_path} ...") + self.algorithm.save_pretrain(self.save_rm_path) + + # finally train the rl agent + self.logger.info("Set up buffer") + self.buffer = self.setup_buffer(self.max_buffer_size, self.buffer_kwargs) + + # self.logger.info(f"Setting up rl datasets and dataloaders ...") + # self._rl_datasets = self.setup_datasets(self.rl_dataset_kwargs) + # if self.rm_label: + # self.logger.info(f"Relabeling the reward using pretrained reward model ...") + # self.algorithm.eval() + # for d in self._rl_datasets: + # d.relabel_reward(self.algorithm) + # self.algorithm.train() + # self._rl_dataloaders, self._rl_dataloaders_iter = self.setup_dataloaders(self._rl_datasets, self.rl_dataloader_kwargs) + + self.logger.info("Warm up") + obs, terminal = self.env.reset(), False + cur_traj_length = 0 + for step in trange(0, self.warmup_step+1): + action = self.env.action_space.sample() + next_obs, reward, terminal, info = self.env.step(action) + # Relabling reward + if self.rm_label: + batch = dict(obs=obs, action=action) + batch = self.algorithm.format_batch(batch) + reward = self.algorithm.select_reward(batch).detach().cpu().numpy() + cur_traj_length += 1 + if cur_traj_length >= self.max_trajectory_length: + terminal = False + self.buffer.add_sample( + { + "obs": obs, + "action": action, + "next_obs": next_obs, + "reward": reward, + "terminal": terminal, + } + ) + obs = next_obs + if terminal or cur_traj_length >= self.max_trajectory_length: + obs = self.env.reset() + cur_traj_length = 0 + + self.logger.info("Start Training") + self.algorithm.train() + + obs, terminal = self.env.reset(), False + cur_traj_length = 0 + traj_return = 0 + true_traj_return = 0 + + for step in trange(0, self.rl_steps+1, desc="RL"): + # random sampling + if step < self.random_policy_step + 1: + action = self.env.action_space.sample() + else: + batch = dict(obs=obs) + batch = self.algorithm.format_batch(batch) + with torch.no_grad(): + action = self.algorithm.select_action(batch) + + next_obs, reward, terminal, info = self.env.step(action) + true_traj_return += reward + cur_traj_length += 1 + if cur_traj_length >= self.max_trajectory_length: + terminal = False + # Relabling reward + if self.rm_label: + batch = dict(obs=obs, action=action) + batch = self.algorithm.format_batch(batch) + reward = self.algorithm.select_reward(batch).detach().cpu().numpy() + traj_return += reward + self.buffer.add_sample( + { + "obs": obs, + "action": action, + "next_obs": next_obs, + "reward": reward, + "terminal": terminal, + } + ) + obs = next_obs + if terminal or cur_traj_length >= self.max_trajectory_length: + obs, terminal = self.env.reset(), False + train_metrics = dict(traj_return=traj_return, true_traj_return=true_traj_return) + self.logger.log_scalars("train", train_metrics, step=step) + cur_traj_length = 0 + traj_return = 0 + true_traj_return = 0 + + batches = self.buffer.random_batch(self.batch_size) + batches = self.algorithm.format_batch(batches) + rl_metrics = self.algorithm.train_step(batches, step=step, total_steps=self.rl_steps) + + if step % self.log_freq == 0: + self.logger.log_scalars("", rl_metrics, step=step) + + if self.eval_freq and step % self.eval_freq == 0: + self.algorithm.eval() + eval_metrics = self.rl_evaluate() + self.logger.log_scalars("eval", eval_metrics, step=step) + self.algorithm.train() + + if self.checkpoint_freq and step % self.checkpoint_freq == 0: + checkpoint_metadata = dict(step=step) + self.algorithm.save(self.logger.output_dir, f"step_{step}.pt", checkpoint_metadata) + + # clean up + checkpoint_metadata = dict(step=step) + self.algorithm.save(self.logger.output_dir, "final.pt", checkpoint_metadata) + if self._env is not None: + self._env.close() + if self._eval_env is not None: + self._eval_env.close() + + def rm_evaluate(self): + if self.rm_eval_kwargs is None: + return {} + if not hasattr(self, "rm_eval_fn"): + self.rm_eval_fn = vars(wiserl.eval)[self.rm_eval_kwargs.pop("function")] + eval_metrics = self.rm_eval_fn( + self.eval_env, self.algorithm, + **self.rm_eval_kwargs + ) + return eval_metrics + + def rl_evaluate(self): + assert not self.algorithm.training + if self.rl_eval_kwargs is None: + return {} + if not hasattr(self, "rl_eval_fn"): + self.rl_eval_fn = vars(wiserl.eval)[self.rl_eval_kwargs.pop("function")] + eval_metrics = self.rl_eval_fn( + self.eval_env, self.algorithm, + **self.rl_eval_kwargs + ) + return eval_metrics