diff --git a/benchmark/benchmark_sharpa_init_dr_construct.py b/benchmark/benchmark_sharpa_init_dr_construct.py index ee557ebca..25510888a 100644 --- a/benchmark/benchmark_sharpa_init_dr_construct.py +++ b/benchmark/benchmark_sharpa_init_dr_construct.py @@ -21,10 +21,10 @@ - `/home/admin1/ws/unilabsim/mujoco_uni/python/mujoco/batch_env.cc` Usage: - uv run python benchmark/benchmark_sharpa_init_dr_construct.py - uv run python benchmark/benchmark_sharpa_init_dr_construct.py --env-nums 256,512,1024 - uv run python benchmark/benchmark_sharpa_init_dr_construct.py --variant-counts 1,2,4,8 - uv run python benchmark/benchmark_sharpa_init_dr_construct.py --measure construct_plus_pool + uv run benchmark/benchmark_sharpa_init_dr_construct.py + uv run benchmark/benchmark_sharpa_init_dr_construct.py --env-nums 256,512,1024 + uv run benchmark/benchmark_sharpa_init_dr_construct.py --variant-counts 1,2,4,8 + uv run benchmark/benchmark_sharpa_init_dr_construct.py --measure construct_plus_pool """ from __future__ import annotations @@ -79,7 +79,7 @@ class ConstructRecord: mode: str variant_count: int num_envs: int - scale_range: list[float] + scale_list: list[float] repeats: int samples_sec: list[float] mean_sec: float @@ -122,9 +122,11 @@ def _parse_variant_counts(value: str | None) -> list[int]: def _compose_cfg(task: str, *, lower: float, upper: float, variant_count: int): config_dir = str(ROOT_DIR / "conf" / "ppo") + scale_list = np.linspace(lower, upper, variant_count, dtype=np.float64) + scale_override = ",".join(f"{float(scale):g}" for scale in scale_list) overrides = [ f"task={task}", - f"env.scale_range=[{lower:g},{upper:g},{variant_count}]", + f"env.domain_rand.scale_list=[{scale_override}]", "hydra.run.dir=.", "hydra.output_subdir=null", "hydra/job_logging=disabled", @@ -254,7 +256,7 @@ def _summarize_record( mode=mode, variant_count=variant_count, num_envs=num_envs, - scale_range=[lower, upper, float(variant_count)], + scale_list=np.linspace(lower, upper, variant_count, dtype=np.float64).tolist(), repeats=len(samples), samples_sec=[float(sample) for sample in samples], mean_sec=float(mean(samples)), diff --git a/conf/appo/config.yaml b/conf/appo/config.yaml index 71e5f3dfb..1e0a956f0 100644 --- a/conf/appo/config.yaml +++ b/conf/appo/config.yaml @@ -61,6 +61,7 @@ training: wandb_notes: null wandb_mode: null sim_backend: mujoco + log_root: null log_dir: null play_only: false no_play: false diff --git a/conf/appo/task/allegro_inhand/motrix.yaml b/conf/appo/task/allegro_inhand/motrix.yaml new file mode 100644 index 000000000..e7ba38ecb --- /dev/null +++ b/conf/appo/task/allegro_inhand/motrix.yaml @@ -0,0 +1,63 @@ +# @package _global_ +training: + task_name: AllegroInhandRotation + sim_backend: motrix + play_steps: 200 + render_spacing: 0.5 + cam_distance: 1.5 + cam_lookat: [0.75, 0.75, 0] + cam_elevation: -20.0 +algo: + num_envs: 16384 + steps_per_env: 8 + max_iterations: 201 + save_interval: 100 + algorithm: + value_loss_coef: 4.0 + entropy_coef: 0.01 + learning_rate: 0.001 + desired_kl: 0.02 + num_learning_epochs: 5 + num_mini_batches: 4 + clip_param: 0.2 + gamma: 0.99 + lam: 0.95 + max_grad_norm: 1.0 + use_clipped_value_loss: true + schedule: adaptive + actor: + hidden_dims: [512, 256, 128] + activation: elu + obs_normalization: true + distribution_cfg: + class_name: rsl_rl.modules.distribution.GaussianDistribution + init_std: 1.0 + std_type: scalar + critic: + hidden_dims: [512, 256, 128] + activation: elu + obs_normalization: true +reward: + scales: + rotate: 1.25 + obj_linvel: -0.3 + pose_diff: -0.3 + torque: -0.1 + work: -2.0 + drop: 0.0 + angvel_clip_min: -0.5 + angvel_clip_max: 0.5 + reset_z_threshold: 0.125 +env: + gen_grasp: false + max_episode_seconds: 20.0 + grasp_cache_path: cache/allegro_grasp_50k.npy + # Keep only grasp/pose reset variation. All online DR terms stay disabled. + domain_rand: + randomize_base_mass: false + random_com: false + randomize_gravity: false + push_robots: false + joint_noise: 0.0 + ball_vel_noise: 0.0 + ball_z_offset: 0.0 diff --git a/conf/appo/task/allegro_inhand/mujoco.yaml b/conf/appo/task/allegro_inhand/mujoco.yaml index 19ebd4df2..008f08d34 100644 --- a/conf/appo/task/allegro_inhand/mujoco.yaml +++ b/conf/appo/task/allegro_inhand/mujoco.yaml @@ -2,16 +2,40 @@ training: task_name: AllegroInhandRotation sim_backend: mujoco + play_steps: 200 + render_spacing: 0.5 + cam_distance: 1.5 + cam_lookat: [0.75, 0.75, 0] + cam_elevation: -20.0 algo: num_envs: 16384 steps_per_env: 8 - max_iterations: 501 + max_iterations: 201 + save_interval: 100 algorithm: value_loss_coef: 4.0 + entropy_coef: 0.01 + learning_rate: 0.001 desired_kl: 0.02 + num_learning_epochs: 5 + num_mini_batches: 4 + clip_param: 0.2 + gamma: 0.99 + lam: 0.95 + max_grad_norm: 1.0 + use_clipped_value_loss: true + schedule: adaptive actor: + hidden_dims: [512, 256, 128] + activation: elu obs_normalization: true + distribution_cfg: + class_name: rsl_rl.modules.distribution.GaussianDistribution + init_std: 1.0 + std_type: scalar critic: + hidden_dims: [512, 256, 128] + activation: elu obs_normalization: true reward: scales: @@ -24,3 +48,16 @@ reward: angvel_clip_min: -0.5 angvel_clip_max: 0.5 reset_z_threshold: 0.125 +env: + gen_grasp: false + max_episode_seconds: 20.0 + grasp_cache_path: cache/allegro_grasp_50k.npy + # Keep only grasp/pose reset variation. All online DR terms stay disabled. + domain_rand: + randomize_base_mass: false + random_com: false + randomize_gravity: false + push_robots: false + joint_noise: 0.0 + ball_vel_noise: 0.0 + ball_z_offset: 0.0 diff --git a/conf/appo/task/sharpa_inhand/mujoco.yaml b/conf/appo/task/sharpa_inhand/mujoco.yaml new file mode 100644 index 000000000..90a06114f --- /dev/null +++ b/conf/appo/task/sharpa_inhand/mujoco.yaml @@ -0,0 +1,116 @@ +# @package _global_ +# Base Sharpa APPO MuJoCo owner config. The HORA variant inherits this file +# so backend support and shared hyperparameters stay visible at the backend +# owner layer rather than being duplicated in a HORA-only variant. + +training: + task_name: SharpaInhandRotation + sim_backend: mujoco + play_steps: 200 + render_spacing: 0.5 + cam_distance: 1.5 + cam_lookat: [0.75, 0.75, 0.4] + cam_elevation: -20.0 + +algo: + num_envs: 16384 + steps_per_env: 8 + max_iterations: 301 + save_interval: 50 + actor: + hidden_dims: [512, 256, 128] + activation: elu + obs_normalization: true + distribution_cfg: + class_name: rsl_rl.modules.distribution.GaussianDistribution + init_std: 1.0 + std_type: scalar + critic: + hidden_dims: [512, 256, 128] + activation: elu + obs_normalization: true + algorithm: + value_loss_coef: 4.0 + entropy_coef: 0.01 + learning_rate: 0.001 + desired_kl: 0.02 + num_learning_epochs: 5 + num_mini_batches: 4 + clip_param: 0.2 + gamma: 0.99 + lam: 0.95 + +reward: + scales: + rotate: 2.5 + obj_linvel: -0.3 + pose_diff: -0.4 + torque: -0.1 + work: -0.5 + object_pos: 0.003 + angvel_clip_min: -0.5 + angvel_clip_max: 0.5 + +env: + zero_action_test_mode: false + clip_obs: 5.0 + clip_actions: 1.0 + reset_height_lower: 0.59906 + reset_height_upper: 0.63906 + reset_angle_diff: 0.7853981633974483 + rot_axis: [0.0, 0.0, 1.0] + grasp_cache_path: cache/sharpa_grasp_linspace + sensor: + tactile_force_sensor_names: + - contact_right_thumb_elastomer_force + - contact_right_index_elastomer_force + - contact_right_middle_elastomer_force + - contact_right_ring_elastomer_force + - contact_right_pinky_elastomer_force + disable_tactile_ids: [] + use_default_object_pose_for_object_pos_anchor: false + obs: + observation_mode: flattened + enable_tactile: true + binary_contact: false + enable_contact_pos: false + contact_smooth: 0.5 + contact_threshold: 0.05 + tactile_force_clip_max: 4.0 + priv_info: + include_friction_scale: true + include_gravity_direction: false + control_config: + action_scale: 0.041666666666666664 + p_gain: 1.0 # use the XML value instead of this value by default + d_gain: 0.1 + torque_control: false # can only be false + dof_limits_scale: 0.9 + domain_rand: + scale_list: [0.8, 0.9, 1.0, 1.1, 1.2, 1.3, 1.4, 1.5, 1.6] + randomize_gravity_direction: true + gravity_direction_magnitude: 9.81 + randomize_pd_gains: true + randomize_p_gain_scale_lower: 0.5 + randomize_p_gain_scale_upper: 2.0 + randomize_d_gain_scale_lower: 0.5 + randomize_d_gain_scale_upper: 2.0 + randomize_friction: true + randomize_friction_scale_lower: 0.75 + randomize_friction_scale_upper: 1.25 + elastomer_base_friction: 2.0 + metal_base_friction: 1.0 + object_base_friction: 2.0 + randomize_com: true + randomize_com_lower: -0.01 + randomize_com_upper: 0.01 + randomize_mass: true + randomize_mass_lower: 0.01 + randomize_mass_upper: 0.25 + force_scale: 2.0 + random_force_prob_scalar: 0.25 + force_decay: 0.9 + force_decay_interval: 0.08 + joint_noise_scale: 0.02 + contact_latency: 0.005 + contact_sensor_noise: 0.01 diff --git a/conf/appo/task/sharpa_inhand/mujoco_hora.yaml b/conf/appo/task/sharpa_inhand/mujoco_hora.yaml new file mode 100644 index 000000000..b86bb01ab --- /dev/null +++ b/conf/appo/task/sharpa_inhand/mujoco_hora.yaml @@ -0,0 +1,32 @@ +# @package _global_ +# HORA Sharpa APPO variant. Inherit the shared MuJoCo owner so backend support +# and shared hyperparameters stay aligned with the non-HORA Sharpa APPO config. +defaults: + - /task/sharpa_inhand/mujoco + - _self_ + +algo: + algo_log_name: hora_appo + runtime_impl: hora_appo + runtime_resolver: unilab.algos.torch.hora.appo:resolve_hora_appo_runtime + obs_groups: + # Keep grouped keys explicit in the owner YAML; runtime support for these + # grouped observations lands in the next implementation step. + actor: + actor: 0 + priv_info: 0 + critic: + actor: 0 + priv_info: 0 + actor: + class_name: unilab.algos.torch.hora:HoraActorModel + priv_info_embed_dim: 9 + priv_mlp_hidden_dims: [256, 128, 9] + critic: + class_name: unilab.algos.torch.hora:HoraCriticModel + priv_info_embed_dim: 9 + priv_mlp_hidden_dims: [256, 128, 9] + +env: + obs: + observation_mode: separated diff --git a/conf/hora_distill/config.yaml b/conf/hora_distill/config.yaml new file mode 100644 index 000000000..2a0a4abcc --- /dev/null +++ b/conf/hora_distill/config.yaml @@ -0,0 +1,44 @@ +defaults: + - _self_ + - task: sharpa_inhand/mujoco + +algo: + algo_log_name: hora_distill + seed: 1 + num_envs: 4096 + max_agent_steps: 1000000000 + save_interval_steps: 100000000 + log_interval_steps: 32768 + learning_rate: 3.0e-4 + load_run: "-1" + checkpoint: -1 + model: {} + +training: + task_name: SharpaInhandRotation + device: null + logger: tensorboard + sim_backend: mujoco + play_only: false + play_env_num: 16 + play_steps: 200 + render_spacing: 1.0 + cam_distance: 6.0 + cam_elevation: -20.0 + cam_azimuth: 90.0 + cam_lookat: null + cam_tracking: false + cam_tracking_env_idx: 0 + cam_tracking_extra_envs: 2 + log_root: null + log_dir: null + +hydra: + run: + dir: . + output_subdir: null + job: + chdir: false + job_logging: + root: + handlers: [console] diff --git a/conf/hora_distill/task/sharpa_inhand/mujoco.yaml b/conf/hora_distill/task/sharpa_inhand/mujoco.yaml new file mode 100644 index 000000000..f3e4b0ab4 --- /dev/null +++ b/conf/hora_distill/task/sharpa_inhand/mujoco.yaml @@ -0,0 +1,19 @@ +# @package _global_ +teacher: + algo_family: ppo + task: sharpa_inhand/mujoco_hora + +training: + task_name: SharpaInhandRotation + sim_backend: mujoco + render_spacing: 0.5 + cam_distance: 1.5 + cam_lookat: [0.75, 0.75, 0.4] + cam_elevation: -20.0 + cam_azimuth: 90.0 + +algo: + algo_log_name: hora_distill + num_envs: 16384 + max_agent_steps: 100000000 + save_interval_steps: 10000000 diff --git a/conf/offpolicy/config.yaml b/conf/offpolicy/config.yaml index 3deee026d..7f9bea99e 100644 --- a/conf/offpolicy/config.yaml +++ b/conf/offpolicy/config.yaml @@ -25,6 +25,7 @@ training: cam_distance: 6.0 cam_elevation: -20.0 cam_azimuth: 90.0 + log_root: null log_dir: null no_sync_collection: false env_steps_per_sync: 1 diff --git a/conf/ppo/config.yaml b/conf/ppo/config.yaml index 3bd69717b..b1d08aafa 100644 --- a/conf/ppo/config.yaml +++ b/conf/ppo/config.yaml @@ -84,6 +84,7 @@ training: cam_tracking: false cam_tracking_env_idx: 0 cam_tracking_extra_envs: 2 + log_root: null num_timesteps: null log_dir: null diff --git a/conf/ppo/config_mlx.yaml b/conf/ppo/config_mlx.yaml index ecf251812..113a26770 100644 --- a/conf/ppo/config_mlx.yaml +++ b/conf/ppo/config_mlx.yaml @@ -72,6 +72,7 @@ training: seed: 1 log_interval: 10 save_interval: 50 + log_root: null hydra: run: diff --git a/conf/ppo/task/allegro_inhand/motrix.yaml b/conf/ppo/task/allegro_inhand/motrix.yaml index 4c791d9a3..f4d4aa130 100644 --- a/conf/ppo/task/allegro_inhand/motrix.yaml +++ b/conf/ppo/task/allegro_inhand/motrix.yaml @@ -6,7 +6,23 @@ algo: num_envs: 16384 num_steps_per_env: 8 max_iterations: 201 - empirical_normalization: true + obs_groups: + actor: [policy] + critic: [policy] + actor: + class_name: rsl_rl.models.MLPModel + hidden_dims: [512, 256, 128] + activation: elu + obs_normalization: true + distribution_cfg: + class_name: rsl_rl.modules.distribution.GaussianDistribution + init_std: 1.0 + std_type: scalar + critic: + class_name: rsl_rl.models.MLPModel + hidden_dims: [512, 256, 128] + activation: elu + obs_normalization: true algorithm: value_loss_coef: 4.0 desired_kl: 0.02 @@ -25,3 +41,12 @@ env: gen_grasp: false max_episode_seconds: 20.0 grasp_cache_path: cache/allegro_grasp_50k.npy + # Keep only grasp/pose reset variation. All online DR terms stay disabled. + domain_rand: + randomize_base_mass: false + random_com: false + randomize_gravity: false + push_robots: false + joint_noise: 0.0 + ball_vel_noise: 0.0 + ball_z_offset: 0.0 diff --git a/conf/ppo/task/allegro_inhand/mujoco.yaml b/conf/ppo/task/allegro_inhand/mujoco.yaml index 9adb4fecb..2de863a8f 100644 --- a/conf/ppo/task/allegro_inhand/mujoco.yaml +++ b/conf/ppo/task/allegro_inhand/mujoco.yaml @@ -2,11 +2,31 @@ training: task_name: AllegroInhandRotation sim_backend: mujoco + render_spacing: 0.5 + cam_distance: 1.5 + cam_lookat: [0.75, 0.75, 0] + cam_elevation: -20.0 algo: num_envs: 16384 num_steps_per_env: 8 max_iterations: 201 - empirical_normalization: true + obs_groups: + actor: [policy] + critic: [policy] + actor: + class_name: rsl_rl.models.MLPModel + hidden_dims: [512, 256, 128] + activation: elu + obs_normalization: true + distribution_cfg: + class_name: rsl_rl.modules.distribution.GaussianDistribution + init_std: 1.0 + std_type: scalar + critic: + class_name: rsl_rl.models.MLPModel + hidden_dims: [512, 256, 128] + activation: elu + obs_normalization: true algorithm: value_loss_coef: 4.0 desired_kl: 0.02 @@ -25,4 +45,12 @@ env: gen_grasp: false max_episode_seconds: 20.0 grasp_cache_path: cache/allegro_grasp_50k.npy - + # Keep only grasp/pose reset variation. All online DR terms stay disabled. + domain_rand: + randomize_base_mass: false + random_com: false + randomize_gravity: false + push_robots: false + joint_noise: 0.0 + ball_vel_noise: 0.0 + ball_z_offset: 0.0 diff --git a/conf/ppo/task/sharpa_inhand/motrix.yaml b/conf/ppo/task/sharpa_inhand/motrix.yaml deleted file mode 100644 index 96e045369..000000000 --- a/conf/ppo/task/sharpa_inhand/motrix.yaml +++ /dev/null @@ -1,89 +0,0 @@ -# @package _global_ -training: - task_name: SharpaInhandRotation - sim_backend: motrix -algo: - num_envs: 16384 - num_steps_per_env: 8 - max_iterations: 2290 - empirical_normalization: true - policy: - actor_hidden_dims: [512, 256, 128] - critic_hidden_dims: [512, 256, 128] - algorithm: - value_loss_coef: 4.0 - entropy_coef: 0.01 - learning_rate: 0.001 - desired_kl: 0.02 - num_learning_epochs: 5 - num_mini_batches: 4 - clip_param: 0.2 - gamma: 0.99 - lam: 0.95 -reward: - scales: - rotate: 2.5 - obj_linvel: -0.3 - pose_diff: -0.4 - torque: -0.1 - work: -0.5 - object_pos: 0.003 - angvel_clip_min: -0.5 - angvel_clip_max: 0.5 -env: - torque_control: false - zero_action_test_mode: false - clip_obs: 5.0 - clip_actions: 1.0 - control_config: - action_scale: 0.041666666666666664 - p_gain: 1.0 - d_gain: 0.1 - reset_height_lower: 0.59906 - reset_height_upper: 0.63906 - reset_angle_diff: 0.7853981633974483 - reset_random_quat: false - rot_axis: [0.0, 0.0, 1.0] - grasp_cache_path: cache/sharpa_grasp_linspace - joint_noise_scale: 0.0 - critic_obs_mode: merged - observation_mode: simple - sensor: - tactile_force_sensor_names: - - contact_right_thumb_elastomer_force - - contact_right_index_elastomer_force - - contact_right_middle_elastomer_force - - contact_right_ring_elastomer_force - - contact_right_pinky_elastomer_force - enable_tactile: true - binary_contact: false - enable_contact_pos: false - disable_tactile_ids: [] - contact_smooth: 0.5 - contact_threshold: 0.05 - contact_latency: 0.005 - contact_sensor_noise: 0.01 - dof_limits_scale: 0.9 - scale_range: [0.5, 0.5, 1] - randomize_pd_gains: false - randomize_p_gain_scale_lower: 0.5 - randomize_p_gain_scale_upper: 2.0 - randomize_d_gain_scale_lower: 0.5 - randomize_d_gain_scale_upper: 2.0 - randomize_friction: false - randomize_friction_scale_lower: 0.5 - randomize_friction_scale_upper: 2.0 - elastomer_base_friction: 0.8 - metal_base_friction: 0.1 - object_base_friction: 0.5 - randomize_com: false - randomize_com_lower: -0.01 - randomize_com_upper: 0.01 - randomize_mass: false - randomize_mass_lower: 0.01 - randomize_mass_upper: 0.25 - force_scale: 0.0 - random_force_prob_scalar: 0.0 - force_decay: 0.9 - force_decay_interval: 0.08 - gravity_curriculum: false diff --git a/conf/ppo/task/sharpa_inhand/mujoco.yaml b/conf/ppo/task/sharpa_inhand/mujoco.yaml index 4f9ceaed1..3f7a202e6 100644 --- a/conf/ppo/task/sharpa_inhand/mujoco.yaml +++ b/conf/ppo/task/sharpa_inhand/mujoco.yaml @@ -9,7 +9,7 @@ training: algo: num_envs: 16384 num_steps_per_env: 8 - max_iterations: 501 + max_iterations: 301 empirical_normalization: true policy: actor_hidden_dims: [512, 256, 128] @@ -36,23 +36,14 @@ reward: angvel_clip_min: -0.5 angvel_clip_max: 0.5 env: - torque_control: false zero_action_test_mode: false clip_obs: 5.0 clip_actions: 1.0 - control_config: - action_scale: 0.041666666666666664 - p_gain: 1.0 - d_gain: 0.1 - reset_height_lower: 0.59906 + reset_height_lower: 0.59906 # assume the hand base height is 0.5 reset_height_upper: 0.63906 reset_angle_diff: 0.7853981633974483 - reset_random_quat: false rot_axis: [0.0, 0.0, 1.0] grasp_cache_path: cache/sharpa_grasp_linspace - joint_noise_scale: 0.0 - critic_obs_mode: merged - observation_mode: simple sensor: tactile_force_sensor_names: - contact_right_thumb_elastomer_force @@ -60,35 +51,50 @@ env: - contact_right_middle_elastomer_force - contact_right_ring_elastomer_force - contact_right_pinky_elastomer_force - enable_tactile: true - binary_contact: false - enable_contact_pos: false disable_tactile_ids: [] - contact_smooth: 0.5 - contact_threshold: 0.05 - contact_latency: 0.005 - contact_sensor_noise: 0.01 - dof_limits_scale: 0.9 - scale_range: [0.8, 1.1, 4] - randomize_pd_gains: false - randomize_p_gain_scale_lower: 0.5 - randomize_p_gain_scale_upper: 2.0 - randomize_d_gain_scale_lower: 0.5 - randomize_d_gain_scale_upper: 2.0 - randomize_friction: false - randomize_friction_scale_lower: 0.5 - randomize_friction_scale_upper: 2.0 - elastomer_base_friction: 0.8 - metal_base_friction: 0.1 - object_base_friction: 0.5 - randomize_com: false - randomize_com_lower: -0.01 - randomize_com_upper: 0.01 - randomize_mass: false - randomize_mass_lower: 0.01 - randomize_mass_upper: 0.25 - force_scale: 0.0 - random_force_prob_scalar: 0.0 - force_decay: 0.9 - force_decay_interval: 0.08 - gravity_curriculum: false + use_default_object_pose_for_object_pos_anchor: false + obs: + observation_mode: flattened # non-hora setting + enable_tactile: true + binary_contact: false + enable_contact_pos: false + contact_smooth: 0.5 + contact_threshold: 0.05 + tactile_force_clip_max: 4.0 + priv_info: + include_friction_scale: true + include_gravity_direction: false + control_config: + action_scale: 0.041666666666666664 + p_gain: 1.0 # use the XML value instead of this value by default + d_gain: 0.1 + torque_control: false # can only be false + dof_limits_scale: 0.9 # tighten the XML dof limits + domain_rand: + scale_list: [0.8, 0.9, 1.0, 1.1, 1.2, 1.3, 1.4, 1.5, 1.6] # cylinder object scales + randomize_gravity_direction: true + gravity_direction_magnitude: 9.81 + randomize_pd_gains: true + randomize_p_gain_scale_lower: 0.5 + randomize_p_gain_scale_upper: 2.0 + randomize_d_gain_scale_lower: 0.5 + randomize_d_gain_scale_upper: 2.0 + randomize_friction: true + randomize_friction_scale_lower: 0.75 + randomize_friction_scale_upper: 1.25 + elastomer_base_friction: 2.0 + metal_base_friction: 1.0 + object_base_friction: 2.0 + randomize_com: true + randomize_com_lower: -0.01 + randomize_com_upper: 0.01 + randomize_mass: true + randomize_mass_lower: 0.01 + randomize_mass_upper: 0.25 + force_scale: 2.0 + random_force_prob_scalar: 0.25 + force_decay: 0.9 + force_decay_interval: 0.08 + joint_noise_scale: 0.02 + contact_latency: 0.005 + contact_sensor_noise: 0.01 diff --git a/conf/ppo/task/sharpa_inhand/mujoco_hora.yaml b/conf/ppo/task/sharpa_inhand/mujoco_hora.yaml new file mode 100644 index 000000000..f22437b3f --- /dev/null +++ b/conf/ppo/task/sharpa_inhand/mujoco_hora.yaml @@ -0,0 +1,36 @@ +# @package _global_ +defaults: + - /task/sharpa_inhand/mujoco + - _self_ + +algo: + algo_log_name: hora_ppo + runtime_impl: hora_ppo + runtime_resolver: unilab.algos.torch.hora.rsl_rl:resolve_hora_ppo_runtime + obs_groups: + actor: [actor] + critic: [actor] + actor: + class_name: unilab.algos.torch.hora:HoraActorModel + hidden_dims: [512, 256, 128] + activation: elu + obs_normalization: true + priv_info_embed_dim: 9 + priv_mlp_hidden_dims: [256, 128, 9] + distribution_cfg: + class_name: GaussianDistribution + init_std: 1.0 + std_type: scalar + critic: + class_name: unilab.algos.torch.hora:HoraCriticModel + hidden_dims: [512, 256, 128] + activation: elu + obs_normalization: true + priv_info_embed_dim: 9 + priv_mlp_hidden_dims: [256, 128, 9] + algorithm: + class_name: unilab.algos.torch.hora:HoraPPO + +env: + obs: + observation_mode: separated diff --git a/conf/ppo/task/sharpa_inhand_grasp/motrix.yaml b/conf/ppo/task/sharpa_inhand_grasp/motrix.yaml deleted file mode 100644 index ceee56789..000000000 --- a/conf/ppo/task/sharpa_inhand_grasp/motrix.yaml +++ /dev/null @@ -1,91 +0,0 @@ -# @package _global_ -training: - task_name: SharpaInhandRotationGrasp - sim_backend: motrix - no_play: true -algo: - num_envs: 16384 - num_steps_per_env: 8 - max_iterations: 2290 - empirical_normalization: true - policy: - actor_hidden_dims: [512, 256, 128] - critic_hidden_dims: [512, 256, 128] - algorithm: - value_loss_coef: 4.0 - entropy_coef: 0.0 - learning_rate: 5.0e-3 - desired_kl: 0.02 - num_learning_epochs: 5 - num_mini_batches: 4 - clip_param: 0.2 - gamma: 0.99 - lam: 0.95 -reward: - scales: - rotate: 0.0 - obj_linvel: 0.0 - pose_diff: 0.0 - torque: 0.0 - work: 0.0 - object_pos: 0.0 - angvel_clip_min: -0.5 - angvel_clip_max: 0.5 -env: - max_episode_seconds: 12.0 - torque_control: false - zero_action_test_mode: false - clip_obs: 5.0 - clip_actions: 1.0 - control_config: - action_scale: 0.041666666666666664 - p_gain: 1.0 - d_gain: 0.1 - reset_height_lower: 0.61406 - reset_height_upper: 0.62406 - reset_angle_diff: 0.5235987755982988 - reset_random_quat: false - rot_axis: [0.0, 0.0, 1.0] - grasp_cache_path: cache/sharpa_grasp_linspace - joint_noise_scale: 0.02 - critic_obs_mode: merged - sensor: - tactile_force_sensor_names: - - contact_right_thumb_elastomer_force - - contact_right_index_elastomer_force - - contact_right_middle_elastomer_force - - contact_right_ring_elastomer_force - - contact_right_pinky_elastomer_force - enable_tactile: true - binary_contact: false - enable_contact_pos: false - disable_tactile_ids: [] - contact_smooth: 0.5 - contact_threshold: 0.05 - contact_latency: 0.005 - contact_sensor_noise: 0.01 - dof_limits_scale: 0.9 - scale_range: [0.5, 0.5, 1] - randomize_pd_gains: false - randomize_p_gain_scale_lower: 0.5 - randomize_p_gain_scale_upper: 2.0 - randomize_d_gain_scale_lower: 0.5 - randomize_d_gain_scale_upper: 2.0 - randomize_friction: false - randomize_friction_scale_lower: 0.5 - randomize_friction_scale_upper: 2.0 - elastomer_base_friction: 0.8 - metal_base_friction: 0.1 - object_base_friction: 0.5 - randomize_com: false - randomize_com_lower: -0.01 - randomize_com_upper: 0.01 - randomize_mass: true - randomize_mass_lower: 0.05 - randomize_mass_upper: 0.051 - force_scale: 0.0 - random_force_prob_scalar: 0.0 - force_decay: 0.9 - force_decay_interval: 0.08 - gravity_curriculum: false - grasp_collection_target: 50000 diff --git a/conf/ppo/task/sharpa_inhand_grasp/mujoco.yaml b/conf/ppo/task/sharpa_inhand_grasp/mujoco.yaml index 88186d01e..b42c120f6 100644 --- a/conf/ppo/task/sharpa_inhand_grasp/mujoco.yaml +++ b/conf/ppo/task/sharpa_inhand_grasp/mujoco.yaml @@ -36,7 +36,38 @@ reward: angvel_clip_max: 0.5 env: max_episode_seconds: 3.0 - torque_control: false + domain_rand: + scale_list: [0.8] + randomize_gravity: false + gravity_range: + - [0.0, 0.0, -9.81] + - [0.0, 0.0, -9.81] + randomize_gravity_direction: false + gravity_direction_magnitude: 9.81 + randomize_pd_gains: false + randomize_p_gain_scale_lower: 0.5 + randomize_p_gain_scale_upper: 2.0 + randomize_d_gain_scale_lower: 0.5 + randomize_d_gain_scale_upper: 2.0 + randomize_friction: false + randomize_friction_scale_lower: 0.5 + randomize_friction_scale_upper: 2.0 + elastomer_base_friction: 0.8 + metal_base_friction: 0.1 + object_base_friction: 0.5 + randomize_com: false + randomize_com_lower: -0.01 + randomize_com_upper: 0.01 + randomize_mass: true + randomize_mass_lower: 0.05 + randomize_mass_upper: 0.051 + force_scale: 0.0 + random_force_prob_scalar: 0.0 + force_decay: 0.9 + force_decay_interval: 0.08 + joint_noise_scale: 0.02 + contact_latency: 0.005 + contact_sensor_noise: 0.01 zero_action_test_mode: false clip_obs: 5.0 clip_actions: 1.0 @@ -44,14 +75,21 @@ env: action_scale: 0.041666666666666664 p_gain: 1.0 d_gain: 0.1 + torque_control: false + dof_limits_scale: 0.9 reset_height_lower: 0.61406 reset_height_upper: 0.62406 reset_angle_diff: 0.5235987755982988 - reset_random_quat: false rot_axis: [0.0, 0.0, 1.0] grasp_cache_path: cache/sharpa_grasp_linspace - joint_noise_scale: 0.02 - critic_obs_mode: merged + obs: + observation_mode: flattened + enable_tactile: true + binary_contact: false + enable_contact_pos: false + contact_smooth: 0.5 + contact_threshold: 0.05 + tactile_force_clip_max: 5.0 sensor: tactile_force_sensor_names: - contact_right_thumb_elastomer_force @@ -59,36 +97,5 @@ env: - contact_right_middle_elastomer_force - contact_right_ring_elastomer_force - contact_right_pinky_elastomer_force - enable_tactile: true - binary_contact: false - enable_contact_pos: false disable_tactile_ids: [] - contact_smooth: 0.5 - contact_threshold: 0.05 - contact_latency: 0.005 - contact_sensor_noise: 0.01 - dof_limits_scale: 0.9 - scale_range: [0.8, 1.1, 4] - randomize_pd_gains: false - randomize_p_gain_scale_lower: 0.5 - randomize_p_gain_scale_upper: 2.0 - randomize_d_gain_scale_lower: 0.5 - randomize_d_gain_scale_upper: 2.0 - randomize_friction: false - randomize_friction_scale_lower: 0.5 - randomize_friction_scale_upper: 2.0 - elastomer_base_friction: 0.8 - metal_base_friction: 0.1 - object_base_friction: 0.5 - randomize_com: false - randomize_com_lower: -0.01 - randomize_com_upper: 0.01 - randomize_mass: true - randomize_mass_lower: 0.05 - randomize_mass_upper: 0.051 - force_scale: 0.0 - random_force_prob_scalar: 0.0 - force_decay: 0.9 - force_decay_interval: 0.08 - gravity_curriculum: false - grasp_collection_target: 50000 + grasp_collection_target: 10000 diff --git a/docs/users/zh_CN/02-simulation-backends.md b/docs/users/zh_CN/02-simulation-backends.md index 3318bbb89..d35d35e30 100644 --- a/docs/users/zh_CN/02-simulation-backends.md +++ b/docs/users/zh_CN/02-simulation-backends.md @@ -49,8 +49,8 @@ uv run scripts/generate_support_matrix.py --write | PPO (torch) | `allegro_inhand` (Allegro in-hand) | Tested | Tested | | PPO (torch) | `allegro_inhand_grasp` (allegro inhand grasp) | Tested | Tested | | PPO (torch) | `go2_handstand` (go2 handstand) | Tested | Tested | -| PPO (torch) | `sharpa_inhand` (sharpa inhand) | Tested | Tested | -| PPO (torch) | `sharpa_inhand_grasp` (sharpa inhand grasp) | Tested | Tested | +| PPO (torch) | `sharpa_inhand` (sharpa inhand) | Tested | - | +| PPO (torch) | `sharpa_inhand_grasp` (sharpa inhand grasp) | Tested | - | | PPO (mlx) | `go1_joystick_flat` (Go1 joystick) | Tested | Tested | | PPO (mlx) | `go2_joystick_flat` (Go2 joystick) | Tested | Tested | | PPO (mlx) | `g1_walk_flat` (G1 walk flat) | Tested | Tested | @@ -59,14 +59,15 @@ uv run scripts/generate_support_matrix.py --write | PPO (mlx) | `allegro_inhand` (Allegro in-hand) | Configured | Configured | | PPO (mlx) | `allegro_inhand_grasp` (allegro inhand grasp) | Configured | Configured | | PPO (mlx) | `go2_handstand` (go2 handstand) | Configured | Configured | -| PPO (mlx) | `sharpa_inhand` (sharpa inhand) | Configured | Configured | -| PPO (mlx) | `sharpa_inhand_grasp` (sharpa inhand grasp) | Configured | Configured | +| PPO (mlx) | `sharpa_inhand` (sharpa inhand) | Configured | - | +| PPO (mlx) | `sharpa_inhand_grasp` (sharpa inhand grasp) | Configured | - | | APPO (torch) | `go1_joystick_flat` (Go1 joystick) | Tested | Registered | | APPO (torch) | `go2_joystick_flat` (Go2 joystick) | Tested | Registered | | APPO (torch) | `g1_walk_flat` (G1 walk flat) | Tested | Registered | | APPO (torch) | `g1_motion_tracking` (G1 motion tracking) | Tested | Tested | | APPO (torch) | `g1_flip_tracking` (G1 flip tracking) | Tested | Tested | -| APPO (torch) | `allegro_inhand` (Allegro in-hand) | Tested | Registered | +| APPO (torch) | `allegro_inhand` (Allegro in-hand) | Tested | Tested | +| APPO (torch) | `sharpa_inhand` (sharpa inhand) | Tested | - | | SAC (torch) | `g1_walk_flat` (G1 walk flat) | Tested | Tested | | SAC (torch) | `g1_walk_rough` (G1 walk rough) | Tested | Registered | | SAC (torch) | `g1_sac_wbt` (g1 sac wbt) | Tested | - | diff --git a/docs/users/zh_CN/06-domain-randomization.md b/docs/users/zh_CN/06-domain-randomization.md index c1f8f211f..ef2b338bb 100644 --- a/docs/users/zh_CN/06-domain-randomization.md +++ b/docs/users/zh_CN/06-domain-randomization.md @@ -34,8 +34,8 @@ | `G1WalkFlat` | 是 | 是:`Domain_Rand + Provider + ResetPlan` | 任务状态采样 + common payload | push | [`g1/joystick.py`](../../../src/unilab/envs/locomotion/g1/joystick.py) | | `G1WalkRough` | 是 | 是:复用 [`G1WalkDomainRandomizationProvider`](../../../src/unilab/envs/locomotion/g1/joystick.py) | 任务状态采样 + common payload | push | [`g1/joystick.py`](../../../src/unilab/envs/locomotion/g1/joystick.py) | | `G1MotionTracking` | 是 | 是:`Domain_Rand + Provider + ResetPlan` | 大量任务特有 reset 采样 + common payload | push | [`motion_tracking/g1/tracking.py`](../../../src/unilab/envs/motion_tracking/g1/tracking.py) | -| `AllegroInhandRotation` | 是 | 是:`DomainRandConfig + Provider + ResetPlan` | 任务特有 reset 采样 + common payload | 无 | [`inhand_rot_allegro/rotation.py`](../../../src/unilab/envs/manipulation/inhand_rot_allegro/rotation.py) | -| `SharpaInhandRotation` | 是 | 是:`InitRandomizationPlan + ResetPlan` | grasp cache 采样 + common payload | 无 | [`sharpa_inhand/rotation.py`](../../../src/unilab/envs/manipulation/sharpa_inhand/rotation.py) | +| `AllegroInhandRotation` | 是 | 是:`DomainRandConfig + Provider + ResetPlan` | 任务特有 reset 采样 + common payload | 无 | [`allegro_inhand/rotation.py`](../../../src/unilab/envs/manipulation/allegro_inhand/rotation.py) | +| `SharpaInhandRotation` | 是 | 是:`InitRandomizationPlan + ResetPlan + IntervalRandomizationPlan` | grasp cache 采样 + common payload | object `body_force` | [`sharpa_inhand/rotation.py`](../../../src/unilab/envs/manipulation/sharpa_inhand/rotation.py) | | `SharpaInhandRotationGrasp` | 是 | 是:复用 Sharpa rotation provider 并覆盖 reset 采样 | grasp collection reset + common payload | 无 | [`sharpa_inhand/grasp_gen.py`](../../../src/unilab/envs/manipulation/sharpa_inhand/grasp_gen.py) | ## 任务域随机化清单 @@ -48,8 +48,8 @@ | `G1WalkRough` | 与 `G1WalkFlat` 相同,直接复用同一个 provider | `push_robots` | kp/kd 默认开启;common payload 和 push 默认关闭 | | `G1MotionTracking` | motion frame 采样;root pose 扰动 `x/y/z/roll/pitch/yaw`;root velocity 扰动 `x/y/z/roll/pitch/yaw`;joint position noise;MuJoCo 下按 joint range clip;`current_actions/last_actions` 清零;可选 `base_mass_delta`;可选 `base_com_offset`;可选 `gravity` | `push_robots` | `pose_randomization`、`velocity_randomization`、`joint_position_range` 默认有非零扰动;common payload 和 push 默认关闭 | | `AllegroInhandRotation` | 若有 grasp cache 则随机采样 grasp;否则对 hand joints 加 `joint_noise`、对球加 `ball_z_offset`;始终对球线速度加 `ball_vel_noise`;可选 common reset randomization payload(含 `gravity`) | 无 | grasp cache 路径可用时默认会采样;`joint_noise`、`ball_vel_noise`、`ball_z_offset` 默认 0;common payload 默认关闭 | -| `SharpaInhandRotation` | grasp cache 按 `scale_ids` 分桶采样;object pose / quat reset;可选 common reset randomization payload(含 `gravity`) | 无 | `scale_range` 默认 `[0.5, 0.5, 1]`,MuJoCo 下会在 init 阶段 materialize object geom scale;common payload 默认关闭 | -| `SharpaInhandRotationGrasp` | hand pose reset;object pose / quat reset;采集成功 grasp 并按 `scale_ids` 分桶保存;可选 `base_mass_delta`;可选 `base_com_offset`;可选 `gravity` | 无 | 默认用于生成 Sharpa grasp cache,cache 文件名包含 `scale_range` tag;common payload 默认关闭 | +| `SharpaInhandRotation` | grasp cache 按 `scale_ids` 分桶采样;object pose / quat reset;可选 common reset randomization payload(含 `gravity`) | object `body_force` direct force disturbance | `domain_rand.scale_list` 默认来自 owner YAML,MuJoCo 下会在 init 阶段 materialize object geom scale;common payload 默认关闭;object force 默认由 Sharpa owner YAML 开启 | +| `SharpaInhandRotationGrasp` | hand pose reset;object pose / quat reset;采集成功 grasp 并按 `scale_ids` 分桶保存;可选 `base_mass_delta`;可选 `base_com_offset`;可选 `gravity` | 无 | 默认用于生成 Sharpa grasp cache,cache 文件名包含单个 scale 值;common payload 默认关闭 | ## 当前统一 DR 能力与边界 @@ -93,9 +93,15 @@ backend capability 当前是: -- [`MuJoCoBackend`](../../../src/unilab/base/backend/mujoco_backend.py):支持上面 7 个 reset term,且支持 interval push +- [`MuJoCoBackend`](../../../src/unilab/base/backend/mujoco_backend.py):支持上面 7 个 reset term,且支持 interval push 与 interval body force - [`MotrixBackend`](../../../src/unilab/base/backend/motrix_backend.py):支持 `base_mass_delta`、`base_com_offset`、`kp`、`kd`,且支持 interval push;初始化阶段要求 actuator 全为 position +注意: + +- 当前 `IntervalRandomizationPlan` 支持 `push_perturbation_limit`、`body_linear_velocity_delta` 与 `body_force`;其中 `body_force` 表达热路径直接外力扰动,不暴露 backend 私有 `xfrc_applied` 细节。 +- 当前 MuJoCo backend 的 interval push 和 interval body force 都通过 `xfrc_applied` 下发外力;Sharpa-hand 的 object disturbance 已切换为 direct force disturbance。 +- Motrix backend 当前仍不支持 direct body-force disturbance,因此此类 owner config 需要继续显式关闭。 + 但任务侧当前实际情况是:并不是所有 provider 都构造这些字段。backend contract 是能力边界,任务配置和 provider 是否下发 payload 才决定该任务是否实际启用对应 DR 项。 ## Reset gravity 用法 @@ -206,41 +212,28 @@ Sharpa-hand 是当前仓库里 `geom_size` init-lifecycle DR 的示例任务。 ### 1. 配置入口 -Sharpa 的缩放配置位于 env owner YAML 的 `env.scale_range`: +Sharpa 的缩放配置位于 env owner YAML 的 `env.domain_rand.scale_list`: ```yaml env: object_body_name: object object_geom_name: object - scale_range: [0.5, 0.8, 4] + domain_rand: + scale_list: [0.5, 0.6, 0.7, 0.8] ``` 字段语义: - `object_body_name`:object body 名称,用于 reset / observation 中定位 object body,不是 scale 的目标字段。 - `object_geom_name`:要缩放的 MuJoCo geom 名称,默认是 `object`。 -- `scale_range[0]`:最小 scale,必须大于 0。 -- `scale_range[1]`:最大 scale,必须大于 0。 -- `scale_range[2]`:scale 桶数量,必须是正整数。 +- `domain_rand.scale_list`:显式 scale 列表;每个值都必须大于 0。 +- `domain_rand.scale_list` 的顺序就是 `scale_id` 顺序。 +- `domain_rand.scale_list` 的长度就是 model variant 数量。 -实际 scale 值由 `np.linspace(lower, upper, num_scales)` 生成。例如: - -```yaml -env: - scale_range: [0.5, 0.8, 4] -``` - -会生成 4 个 scale: - -- `0.5` -- `0.6` -- `0.7` -- `0.8` - -每个 env 会被静态分配一个 `scale_id`。当前分配规则是按 bucket 连续分配,因此 `algo.num_envs` 必须能被 `num_scales` 整除: +每个 env 会被静态分配一个 `scale_id`。当前分配规则是按 bucket 连续分配;当 `algo.num_envs` 不能被 `num_scales` 整除时,前几个 scale bucket 会多分配一个 env: ```bash -uv run scripts/train_rsl_rl.py task=sharpa_inhand/mujoco 'env.scale_range=[0.5,0.8,4]' algo.num_envs=4096 +uv run scripts/train_rsl_rl.py task=sharpa_inhand/mujoco 'env.domain_rand.scale_list=[0.5,0.6,0.7,0.8]' algo.num_envs=4096 ``` 如果 `algo.num_envs=4096`、`num_scales=4`,则每 1024 个 env 使用同一个 scale bucket。 @@ -249,46 +242,55 @@ uv run scripts/train_rsl_rl.py task=sharpa_inhand/mujoco 'env.scale_range=[0.5,0 MuJoCo backend 的落地方式是: -1. env/provider 在 init 阶段根据 `scale_range` 构造 `ModelVariantSpec`。 +1. env/provider 在 init 阶段根据 `scale_list` 构造 `ModelVariantSpec`。 2. backend 用 `MjSpec` 读取模型并修改 `object_geom_name` 对应 geom 的 `size`。 3. 每个 scale 编译一套 scale-specific `MjModel`。 4. 第一次需要 physics pool 时,用 env-to-model assignment 展开成长度为 `num_envs` 的 model sequence,再构造 `BatchEnvPool`。 -因此,`scale_range` 只在 env/backend 初始化阶段生效。env 创建后再改 `env.scale_range` 不会改变已经 materialize 的模型池。 +因此,`domain_rand.scale_list` 只在 env/backend 初始化阶段生效。env 创建后再改 `env.domain_rand.scale_list` 不会改变已经 materialize 的模型池。 这个流程有三个重要边界: -- `BatchEnvPool` 是 lazy 构造的;正常路径不会先为默认模型构造一套 pool,再为了 `scale_range` 重建一套 pool。 +- `BatchEnvPool` 是 lazy 构造的;正常路径不会先为默认模型构造一套 pool,再为了 `scale_list` 重建一套 pool。 - 多个 model variant 的编译使用 process-based parallelism 分块执行;不要在 Python thread 里编译,也不要在上层 for 循环串行编译 `num_envs` 个模型。 - worker 用 `MjSpec` 编译 variant 并保存 `.mjb`,父进程只按 `.mjb` 路径加载 `MjModel.from_binary_path(...)`;不要通过 IPC 回传修改后的模型对象或模型 bytes。 ### 3. grasp cache 与 scale bucket -Sharpa rotation 任务按 `scale_ids` 从 grasp cache 里分桶采样: +Sharpa rotation 任务按 `scale_ids` 从多份单-scale grasp cache 采样: + +- cache 文件名默认由 `grasp_cache_path` 和单个 scale 值共同决定。 +- `scale_list: [0.5, 0.6, 0.7, 0.8]` 默认对应 `cache/sharpa_grasp_linspace_0.5.npy`、`cache/sharpa_grasp_linspace_0.6.npy`、`cache/sharpa_grasp_linspace_0.7.npy`、`cache/sharpa_grasp_linspace_0.8.npy`。 +- rotation 启动时会检查 `scale_list` 对应的所有 cache 文件是否存在,缺失即报错。 +- 每个 scale bucket 只从对应 scale 的 cache 文件采样,避免不同 object scale 混用 grasp 初始状态。 -- cache 文件名默认由 `grasp_cache_path` 和 `scale_range` 共同决定。 -- `scale_range: [0.5, 0.8, 4]` 默认对应类似 `cache/sharpa_grasp_linspace_0.5-0.8-4.npy` 的路径。 -- cache 行数必须能被 `num_scales` 整除。 -- 每个 scale bucket 使用自己的 cache 区间,避免不同 object scale 混用 grasp 初始状态。 +生成多 scale cache 时,应分别跑多次 grasp 采集任务: + +```bash +uv run scripts/train_rsl_rl.py task=sharpa_inhand_grasp/mujoco 'env.domain_rand.scale_list=[0.5]' algo.num_envs=4096 +uv run scripts/train_rsl_rl.py task=sharpa_inhand_grasp/mujoco 'env.domain_rand.scale_list=[0.6]' algo.num_envs=4096 +uv run scripts/train_rsl_rl.py task=sharpa_inhand_grasp/mujoco 'env.domain_rand.scale_list=[0.7]' algo.num_envs=4096 +uv run scripts/train_rsl_rl.py task=sharpa_inhand_grasp/mujoco 'env.domain_rand.scale_list=[0.8]' algo.num_envs=4096 +``` -生成多 scale cache 时,应使用同一套 `scale_range` 跑 grasp 采集任务: +也可以顺序执行仓库里的 helper: ```bash -uv run scripts/train_rsl_rl.py task=sharpa_inhand_grasp/mujoco 'env.scale_range=[0.5,0.8,4]' algo.num_envs=4096 +./scripts/sharpa_collect_grasps.sh 0.5 0.6 0.7 0.8 ``` -随后训练 rotation 时使用相同的 `scale_range`: +随后训练 rotation 时使用相同的 `scale_list`: ```bash -uv run scripts/train_rsl_rl.py task=sharpa_inhand/mujoco 'env.scale_range=[0.5,0.8,4]' algo.num_envs=4096 +uv run scripts/train_rsl_rl.py task=sharpa_inhand/mujoco 'env.domain_rand.scale_list=[0.5,0.6,0.7,0.8]' algo.num_envs=4096 ``` ### 4. 边界和注意事项 - `geom_size` 不是 reset DR 字段,不能写进 `ResetPlan.randomization`。 - `BatchEnvPool.reset(..., randomization=...)` 当前不支持 `geom_size`。 -- `geom_size` scale 只在 MuJoCo backend 下 materialize;Motrix backend 当前不会按 `scale_range` 生成多模型池。 -- `scale_range[2]` 是模型 variant 数量,不是每次 reset 随机抽样次数。 +- `geom_size` scale 只在 MuJoCo backend 下 materialize;Motrix backend 当前不会按 `scale_list` 生成多模型池。 +- `scale_list` 的长度是模型 variant 数量,不是每次 reset 随机抽样次数。 - 每个 env 的 `scale_id` 是 init 阶段静态 assignment,不会在 reset 时变化。 - 扩 scale 时应扩 model variant 数量,不应按 `num_envs` 编译一 env 一模型;多个 env 共享同一个 scale bucket 对应的 `MjModel`。 - 热路径不得读取 XML、解析 asset 或用 `getattr` / `hasattr` 探测 backend 私有能力来决定 scale 行为。 diff --git a/docs/users/zh_CN/07-dexterous-inhand-manipulation.md b/docs/users/zh_CN/07-dexterous-inhand-manipulation.md new file mode 100644 index 000000000..890e63904 --- /dev/null +++ b/docs/users/zh_CN/07-dexterous-inhand-manipulation.md @@ -0,0 +1,218 @@ +# Dexterous In-Hand Manipulation 训练 + +语言: 简体中文 + +本页只说明如何运行当前仓库已有的 dexterous inhand manipulation 流程。后端选择必须通过 `task=/` 完成,不要单独 override `training.sim_backend` 来切后端。 + +## Allegro Inhand Rotation + +Allegro 的环境注册名是 `AllegroInhandRotation`,常规训练 task owner 是 `allegro_inhand`。完整流程是先生成 grasp cache,再训练 rotation policy。 + +`allegro_inhand` 是一个 in-hand manipulation 的最小训练示例。策略观测包含 privileged information,默认不启用 domain randomization。 + +### 配置文件 + +- `scripts/train_rsl_rl.py` 主配置:`conf/ppo/config.yaml` +- `task=allegro_inhand/mujoco`:`conf/ppo/task/allegro_inhand/mujoco.yaml` +- `task=allegro_inhand/motrix`:`conf/ppo/task/allegro_inhand/motrix.yaml` +- `task=allegro_inhand_grasp/mujoco`:`conf/ppo/task/allegro_inhand_grasp/mujoco.yaml`,并继承 `conf/ppo/task/allegro_inhand/mujoco.yaml` +- `task=allegro_inhand_grasp/motrix`:`conf/ppo/task/allegro_inhand_grasp/motrix.yaml`,并继承 `conf/ppo/task/allegro_inhand/motrix.yaml` +- `scripts/train_appo.py` 主配置:`conf/appo/config.yaml` +- `task=allegro_inhand/mujoco`:`conf/appo/task/allegro_inhand/mujoco.yaml` +- `task=allegro_inhand/motrix`:`conf/appo/task/allegro_inhand/motrix.yaml` + +### 1. 生成 Grasp Cache + +grasp cache 生成任务是 `allegro_inhand_grasp`: + +```bash +uv run scripts/train_rsl_rl.py task=allegro_inhand_grasp/mujoco training.no_play=true +``` + +Motrix owner 也存在: + +```bash +uv run scripts/train_rsl_rl.py task=allegro_inhand_grasp/motrix training.no_play=true +``` + +默认 rotation 配置读取: + +```text +cache/allegro_grasp_50k.npy +``` + +如果使用自定义 cache,在训练时指定: + +```bash +uv run scripts/train_rsl_rl.py \ + task=allegro_inhand/mujoco \ + env.grasp_cache_path=cache/my_allegro_grasp.npy +``` + +### 2. 训练 Policy + +PPO: + +```bash +uv run scripts/train_rsl_rl.py task=allegro_inhand/mujoco +uv run scripts/train_rsl_rl.py task=allegro_inhand/motrix +``` + +APPO: + +```bash +uv run scripts/train_appo.py task=allegro_inhand/mujoco +uv run scripts/train_appo.py task=allegro_inhand/motrix +``` + +### 3. 回放 + +MuJoCo 会导出 `play_video.mp4`: + +```bash +uv run scripts/train_rsl_rl.py task=allegro_inhand/mujoco training.play_only=true +uv run scripts/train_appo.py task=allegro_inhand/mujoco training.play_only=true +``` + +Motrix 当前使用原生交互式 renderer: + +```bash +uv run scripts/train_rsl_rl.py task=allegro_inhand/motrix training.play_only=true +uv run scripts/train_appo.py task=allegro_inhand/motrix training.play_only=true +``` + +macOS / MacBook 上如果会打开 MotrixSim 原生 renderer,使用 `mxpython`: + +```bash +uv run mxpython scripts/train_rsl_rl.py task=allegro_inhand/motrix training.play_only=true +``` + +## Sharpa Inhand Rotation + +Sharpa 的环境注册名是 `SharpaInhandRotation`,常规训练 task owner 是 `sharpa_inhand`。当前训练 pipeline 是生成 grasp cache、训练 teacher policy、再训练 student policy。 + +`sharpa_inhand` 是一个完整的 [HORA](https://github.com/HaozhiQi/hora) 风格训练示例,训练流程包含完整的 domain randomization。 + +`sharpa_inhand` 当前只支持 MuJoCo backend,不支持 Motrix,因为 Motrix 目前还不支持完整的 domain randomization。 + +### 配置文件 + +- `scripts/train_rsl_rl.py` 主配置:`conf/ppo/config.yaml` +- `task=sharpa_inhand_grasp/mujoco`:`conf/ppo/task/sharpa_inhand_grasp/mujoco.yaml` +- `task=sharpa_inhand/mujoco_hora`:`conf/ppo/task/sharpa_inhand/mujoco_hora.yaml`,并继承 `conf/ppo/task/sharpa_inhand/mujoco.yaml` +- `scripts/train_appo.py` 主配置:`conf/appo/config.yaml` +- `task=sharpa_inhand/mujoco_hora`:`conf/appo/task/sharpa_inhand/mujoco_hora.yaml`,并继承 `conf/appo/task/sharpa_inhand/mujoco.yaml` +- `scripts/train_hora_distill.py` 主配置:`conf/hora_distill/config.yaml` +- student `task=sharpa_inhand/mujoco`:`conf/hora_distill/task/sharpa_inhand/mujoco.yaml` +- student 默认 `teacher.algo_family=ppo`,`teacher.task=sharpa_inhand/mujoco_hora` 指向 `conf/ppo/task/sharpa_inhand/mujoco_hora.yaml` +- 如果蒸馏 APPO teacher,保持 `teacher.task=sharpa_inhand/mujoco_hora`,并将 `teacher.algo_family=appo`,此时会解析到 `conf/appo/task/sharpa_inhand/mujoco_hora.yaml` + +### 1. 生成 Grasp Cache + +grasp cache 生成任务是 `sharpa_inhand_grasp`. 按 object scale 分别采集: + +```bash +uv run scripts/train_rsl_rl.py task=sharpa_inhand_grasp/mujoco 'env.domain_rand.scale_list=[0.8]' training.no_play=true +uv run scripts/train_rsl_rl.py task=sharpa_inhand_grasp/mujoco 'env.domain_rand.scale_list=[1.0]' training.no_play=true +uv run scripts/train_rsl_rl.py task=sharpa_inhand_grasp/mujoco 'env.domain_rand.scale_list=[1.2]' training.no_play=true +``` + +也可以用批量脚本一次生成多组 scale 的 grasp cache: + +```bash +bash scripts/sharpa_collect_grasps.sh 0.8 0.9 1.0 1.1 1.2 1.3 1.4 1.5 1.6 +``` + +常规 Sharpa rotation 配置读取: + +```text +cache/sharpa_grasp_linspace +``` + +如果使用自定义 cache,在后续训练命令中指定: + +```bash +uv run scripts/train_rsl_rl.py \ + task=sharpa_inhand/mujoco \ + env.grasp_cache_path=cache/my_sharpa_grasp_cache +``` + +### 2. 训练 Teacher Policy + +PPO teacher: + +```bash +uv run scripts/train_rsl_rl.py task=sharpa_inhand/mujoco_hora +``` + +APPO teacher: + +```bash +uv run scripts/train_appo.py task=sharpa_inhand/mujoco_hora +``` + +### 3. 训练 Student Policy + +从 PPO teacher 蒸馏 student: + +```bash +uv run scripts/train_hora_distill.py task=sharpa_inhand/mujoco +``` + +从 APPO teacher 蒸馏 student: + +```bash +uv run scripts/train_hora_distill.py \ + task=sharpa_inhand/mujoco \ + teacher.algo_family=appo \ + teacher.task=sharpa_inhand/mujoco_hora +``` + +指定 teacher run: + +```bash +uv run scripts/train_hora_distill.py \ + task=sharpa_inhand/mujoco \ + algo.load_run="2026-04-28_12-00-00_mujoco" +``` + +指定 APPO teacher run: + +```bash +uv run scripts/train_hora_distill.py \ + task=sharpa_inhand/mujoco \ + teacher.algo_family=appo \ + teacher.task=sharpa_inhand/mujoco_hora \ + algo.load_run="2026-04-28_12-00-00_mujoco" +``` + +### 4. 回放 + +回放 teacher: + +```bash +uv run scripts/train_rsl_rl.py task=sharpa_inhand/mujoco_hora training.play_only=true +uv run scripts/train_appo.py task=sharpa_inhand/mujoco_hora training.play_only=true +``` + +回放 student: + +```bash +uv run scripts/train_hora_distill.py task=sharpa_inhand/mujoco training.play_only=true +``` + +## 常用命令 + +日志目录按 `algo.algo_log_name` 和环境名分组,常见路径包括: + +- `logs/rsl_rl_ppo/AllegroInhandRotation//` +- `logs/appo/AllegroInhandRotation//` +- `logs/hora_ppo/SharpaInhandRotation//` +- `logs/hora_appo/SharpaInhandRotation//` +- `logs/hora_distill/SharpaInhandRotation//` + +## Navigation + +- Index: [Documentation](../../README.md) +- Previous: [Domain Randomization](06-domain-randomization.md) +- Next: [Simulation Backends](02-simulation-backends.md) diff --git a/scripts/sharpa_collect_grasps.sh b/scripts/sharpa_collect_grasps.sh new file mode 100755 index 000000000..1d31ac2c1 --- /dev/null +++ b/scripts/sharpa_collect_grasps.sh @@ -0,0 +1,16 @@ +#!/usr/bin/env bash + +set -euo pipefail + +if [ "$#" -lt 1 ]; then + echo "Usage: $0 [scale2 ...]" + exit 1 +fi + +for scale in "$@"; do + echo "[sharpa_collect_grasps] collecting scale=${scale}" + + uv run scripts/train_rsl_rl.py \ + task=sharpa_inhand_grasp/mujoco \ + "env.domain_rand.scale_list=[${scale}]" +done diff --git a/scripts/train_appo.py b/scripts/train_appo.py index 5664f3fb6..6d190776b 100644 --- a/scripts/train_appo.py +++ b/scripts/train_appo.py @@ -5,6 +5,7 @@ import datetime import os import sys +from collections.abc import Callable from pathlib import Path from typing import Any, cast @@ -15,6 +16,7 @@ ROOT_DIR = Path(__file__).parent.parent sys.path.append(str(ROOT_DIR)) +from unilab.algos.torch.appo.runtime import resolve_appo_runtime from unilab.training import ( BackendAdapter, create_env, @@ -92,11 +94,15 @@ def run_motrix_play_loop( def resolve_appo_checkpoint_path( base_log_dir: str | Path, - load_run: str, + load_run: str | int, ) -> tuple[str | None, str | None]: from unilab.training import resolve_checkpoint_path - checkpoint_path, checkpoint_dir = resolve_checkpoint_path(base_log_dir, load_run, suffix=".pt") + checkpoint_path, checkpoint_dir = resolve_checkpoint_path( + base_log_dir, + str(load_run), + suffix=".pt", + ) return ( str(checkpoint_path) if checkpoint_path is not None else None, str(checkpoint_dir) if checkpoint_dir is not None else None, @@ -107,8 +113,29 @@ def _get_log_root(cfg: DictConfig) -> str: return str(get_log_root(ROOT_DIR, cfg)) -def play_appo(cfg: DictConfig, rl_cfg: dict[str, Any]) -> str | None: - """Play mode for APPO.""" +def play_appo( + cfg: DictConfig, + rl_cfg: dict[str, Any], + *, + root_dir: Path | None = None, + resolve_checkpoint_path: Callable[[DictConfig], tuple[str | None, str | None]] | None = None, +) -> str | None: + """Play mode for the default APPO runtime. + + Args: + cfg: Resolved Hydra config for the current run. + rl_cfg: Resolved algorithm config dictionary from Hydra composition. + root_dir: Optional project root forwarded by generic runtime callers. + The default APPO runtime does not need it and ignores the value. + resolve_checkpoint_path: Optional checkpoint resolver injected by the + generic script. When omitted, this function falls back to the + default log-root based APPO checkpoint resolution. + + Returns: + Output video path for offscreen rendering, or ``None`` when running the + native Motrix viewer or when no checkpoint could be resolved. + """ + del root_dir import numpy as np from rsl_rl.utils import resolve_callable from tensordict import TensorDict @@ -174,9 +201,12 @@ def play_appo(cfg: DictConfig, rl_cfg: dict[str, Any]) -> str | None: actor = actor.to(device) actor.eval() - log_root = _get_log_root(cfg) - base_log_dir = os.path.join(log_root, cfg.training.task_name) - load_path, load_path_dir = resolve_appo_checkpoint_path(base_log_dir, cfg.algo.load_run) + if resolve_checkpoint_path is not None: + load_path, load_path_dir = resolve_checkpoint_path(cfg) + else: + log_root = _get_log_root(cfg) + base_log_dir = os.path.join(log_root, cfg.training.task_name) + load_path, load_path_dir = resolve_appo_checkpoint_path(base_log_dir, cfg.algo.load_run) if not load_path or not os.path.exists(load_path): print(f"Could not find run to load. load_path={load_path}") @@ -261,6 +291,9 @@ def forward(self, obs: torch.Tensor) -> torch.Tensor: render_play_mode( env, sim_backend=cfg.training.sim_backend, + render_spacing=float( + getattr(cfg.training, "render_spacing", getattr(env.cfg, "render_spacing", 1.0)) + ), num_steps=num_steps, output_video=output_video, initialize=lambda: np.asarray( @@ -285,6 +318,7 @@ def forward(self, obs: torch.Tensor) -> torch.Tensor: "cam_distance": cfg.training.cam_distance, "cam_elevation": cfg.training.cam_elevation, "cam_azimuth": cfg.training.cam_azimuth, + "cam_lookat": getattr(cfg.training, "cam_lookat", None), "cam_tracking": getattr(cfg.training, "cam_tracking", False), "cam_tracking_env_idx": getattr(cfg.training, "cam_tracking_env_idx", 0), "cam_tracking_extra_envs": getattr(cfg.training, "cam_tracking_extra_envs", 2), @@ -308,6 +342,7 @@ def main(cfg: DictConfig) -> None: if not isinstance(rl_cfg_raw, dict): raise TypeError("cfg.algo must resolve to a dict") rl_cfg = cast(dict[str, Any], rl_cfg_raw) + appo_runtime = resolve_appo_runtime(rl_cfg, default_play_fn=play_appo) if cfg.training.log_dir is None: timestamp = datetime.datetime.now().strftime("%Y-%m-%d_%H-%M-%S") @@ -349,9 +384,7 @@ def main(cfg: DictConfig) -> None: try: if not cfg.training.play_only: - from unilab.algos.torch.appo.runner import APPORunner - - runner = APPORunner( + runner = appo_runtime.runner_cls( **build_appo_runner_kwargs( cfg, env_cfg_override=env_cfg_override, @@ -373,7 +406,15 @@ def main(cfg: DictConfig) -> None: runner.close() if cfg.training.play_only or not cfg.training.no_play: - play_video_path = play_appo(cfg, rl_cfg) + play_video_path = appo_runtime.play_fn( + cfg, + rl_cfg, + root_dir=ROOT_DIR, + resolve_checkpoint_path=lambda current_cfg: resolve_appo_checkpoint_path( + os.path.join(_get_log_root(current_cfg), current_cfg.training.task_name), + current_cfg.algo.load_run, + ), + ) if tracker is not None: tracker.log_video(play_video_path) finally: diff --git a/scripts/train_hora_distill.py b/scripts/train_hora_distill.py new file mode 100644 index 000000000..42368046a --- /dev/null +++ b/scripts/train_hora_distill.py @@ -0,0 +1,362 @@ +import datetime +import json +import sys +from pathlib import Path +from typing import Any, cast + +import hydra +import torch +from omegaconf import DictConfig, OmegaConf +from tensordict import TensorDict + +ROOT_DIR = Path(__file__).parent.parent +SRC_DIR = ROOT_DIR / "src" +if str(SRC_DIR) not in sys.path: + sys.path.insert(0, str(SRC_DIR)) +if str(ROOT_DIR) not in sys.path: + sys.path.insert(0, str(ROOT_DIR)) + +from unilab.algos.torch.hora import HoraDistillationTrainer +from unilab.algos.torch.hora.distill import ( + build_student_actor_and_normalizer, + load_distilled_checkpoint, +) +from unilab.algos.torch.hora.distill_config import ( + apply_teacher_defaults as _apply_teacher_defaults, +) +from unilab.algos.torch.hora.distill_config import ( + get_teacher_owner_spec as _get_teacher_owner_spec, +) +from unilab.algos.torch.hora.distill_config import ( + resolve_teacher_checkpoint_path as _resolve_teacher_checkpoint_path, +) +from unilab.algos.torch.hora.distill_config import ( + resolved_distill_runtime_cfg as _resolved_distill_runtime_cfg, +) +from unilab.algos.torch.hora.distill_config import ( + teacher_run_metadata as _teacher_run_metadata, +) +from unilab.algos.torch.hora.rsl_rl import HoraRslRlVecEnvWrapper as RslRlVecEnvWrapper +from unilab.base.backend.xml import materialize_scene_visual_override +from unilab.training import ( + BackendAdapter, + create_env, + ensure_registries, + get_latest_run, + get_log_root, + setup_logger, +) +from unilab.visualization import render_play_mode + + +def _write_distill_run_config( + log_dir: Path, + *, + cfg: DictConfig, + teacher_metadata: dict[str, Any], +) -> None: + """Persist distillation run config plus teacher provenance near the checkpoints. + + Args: + log_dir: Run directory where the metadata file should be written. + cfg: Resolved distillation config for this run. + teacher_metadata: Explicit teacher provenance dictionary for this run. + + Returns: + None. Writes `distill_run_config.json` into `log_dir`. + """ + payload = { + "run": { + "algo": "hora_distill", + "task": str(OmegaConf.select(cfg, "training.task_name")), + "sim_backend": str(OmegaConf.select(cfg, "training.sim_backend")), + "log_dir": str(log_dir), + "teacher": teacher_metadata, + }, + "config": OmegaConf.to_container(cfg, resolve=True), + } + with (log_dir / "distill_run_config.json").open("w", encoding="utf-8") as f: + json.dump(payload, f, indent=2, ensure_ascii=True) + f.write("\n") + + +def _build_env_cfg_override(cfg: DictConfig) -> dict[str, Any]: + adapter = BackendAdapter( + cfg, + root_dir=ROOT_DIR, + algo_name="hora_distill", + scene_materializer=materialize_scene_visual_override, + ) + return cast(dict[str, Any], adapter.build_task_env_cfg_override()) + + +def _build_play_env_cfg_override(cfg: DictConfig) -> dict[str, Any]: + adapter = BackendAdapter( + cfg, + root_dir=ROOT_DIR, + algo_name="hora_distill", + scene_materializer=materialize_scene_visual_override, + ) + return cast(dict[str, Any], adapter.build_play_env_cfg_override()) + + +def _resolve_stage2_checkpoint_path(cfg: DictConfig) -> tuple[Path | None, Path | None]: + task_log_root = get_log_root(ROOT_DIR, cfg) / str(cfg.training.task_name) + load_run = str(OmegaConf.select(cfg, "algo.load_run", default="-1")) + selected_checkpoint = OmegaConf.select(cfg, "algo.checkpoint", default=-1) + + run_dir: Path | None + if load_run == "-1": + run_dir = get_latest_run(task_log_root) + else: + candidate = Path(load_run) + if not candidate.exists(): + candidate = task_log_root / load_run + if candidate.is_file(): + return candidate, candidate.parent + run_dir = candidate if candidate.is_dir() else None + + if run_dir is None: + return None, None + + if selected_checkpoint not in (None, "", -1, "-1"): + checkpoint_name = ( + f"hora_stage2_{selected_checkpoint}.pt" + if str(selected_checkpoint).isdigit() + else str(selected_checkpoint) + ) + checkpoint_path = run_dir / checkpoint_name + return (checkpoint_path, run_dir) if checkpoint_path.exists() else (None, run_dir) + + last_path = run_dir / "hora_stage2_last.pt" + if last_path.exists(): + return last_path, run_dir + + numbered = [ + path for path in run_dir.glob("hora_stage2_*.pt") if path.stem.split("_")[-1].isdigit() + ] + if not numbered: + return None, run_dir + return max(numbered, key=lambda path: int(path.stem.split("_")[-1])), run_dir + + +def _format_stage2_play_checkpoint_error( + cfg: DictConfig, + *, + task_log_root: Path, + load_path: Path | None, + load_path_dir: Path | None, +) -> str: + selected_checkpoint = OmegaConf.select(cfg, "algo.checkpoint", default=-1) + checkpoint_hint = ( + f" algo.checkpoint={selected_checkpoint!r}" + if selected_checkpoint not in (None, "", -1, "-1") + else "" + ) + if load_path_dir is not None and load_path is None and checkpoint_hint: + reason = f"Requested stage-2 checkpoint was not found under resolved_run={load_path_dir}." + elif not task_log_root.exists(): + reason = "Task log root does not exist." + else: + latest_run = get_latest_run(task_log_root) + if latest_run is None: + reason = "No run directories were found under the task log root." + else: + reason = "Requested run or stage-2 checkpoint could not be resolved." + return ( + "Could not resolve a stage-2 HORA checkpoint for play mode. " + f"{reason} task={cfg.training.task_name} task_log_root={task_log_root} " + f"algo.load_run={cfg.algo.load_run!r}{checkpoint_hint}. " + "Use algo.load_run= and optionally " + "algo.checkpoint=." + ) + + +def _student_policy( + actor, hist_normalizer, obs: TensorDict, *, device: torch.device +) -> torch.Tensor: + proprio_hist = hist_normalizer(obs["proprio_hist"].to(device), update=False) + policy_obs = TensorDict( + { + "actor": obs["actor"].to(device), + "proprio_hist": proprio_hist, + }, + batch_size=obs.batch_size, + device=device, + ) + return actor(policy_obs, stochastic_output=False).clamp_(-1.0, 1.0) + + +def _play_camera_kwargs(cfg: DictConfig) -> dict[str, Any]: + camera_kwargs = { + "cam_tracking": getattr(cfg.training, "cam_tracking", False), + "cam_tracking_env_idx": getattr(cfg.training, "cam_tracking_env_idx", 0), + "cam_tracking_extra_envs": getattr(cfg.training, "cam_tracking_extra_envs", 2), + } + for key in ("cam_distance", "cam_elevation", "cam_azimuth", "cam_lookat"): + value = getattr(cfg.training, key, None) + if value is not None: + camera_kwargs[key] = value + return camera_kwargs + + +def _cfg_with_checkpoint_runtime(cfg: DictConfig, checkpoint: dict[str, Any]) -> DictConfig: + """Merge teacher-independent runtime config stored in a stage-2 checkpoint. + + Args: + cfg: Hydra-composed distillation config supplied to play mode. + checkpoint: Loaded stage-2 checkpoint dictionary. + + Returns: + Config with checkpoint runtime fields restored for environment and model construction. + """ + runtime_cfg = checkpoint.get("distill_runtime_cfg") + if runtime_cfg is None: + # Backward compatibility for older stage-2 checkpoints that did not + # persist teacher-independent playback config. + return _apply_teacher_defaults(cfg) + # Hydra keeps the distillation root config structured, but runtime playback + # metadata legitimately restores owner fields such as reward/env that are + # absent from the bare distillation config. + cfg_clone = OmegaConf.create(OmegaConf.to_container(cfg, resolve=False)) + return cast(DictConfig, OmegaConf.merge(cfg_clone, OmegaConf.create(runtime_cfg))) + + +def play_hora_distill(cfg: DictConfig, device: str) -> str | None: + task_log_root = get_log_root(ROOT_DIR, cfg) / str(cfg.training.task_name) + load_path, load_path_dir = _resolve_stage2_checkpoint_path(cfg) + if load_path is None or load_path_dir is None or not load_path.exists(): + print( + _format_stage2_play_checkpoint_error( + cfg, + task_log_root=task_log_root, + load_path=load_path, + load_path_dir=load_path_dir, + ) + ) + return None + + print(f"Loading distilled model: {load_path}") + checkpoint = torch.load(load_path, map_location="cpu", weights_only=False) + if "model_state_dict" not in checkpoint: + print( + f"Checkpoint at {load_path} is not a HORA distillation checkpoint " + f"(found keys: {set(checkpoint.keys())}). Aborting play." + ) + return None + + cfg = _cfg_with_checkpoint_runtime(cfg, checkpoint) + env = create_env( + cfg, + num_envs=int(cfg.training.play_env_num), + env_cfg_override=_build_play_env_cfg_override(cfg), + ) + wrapped_env = RslRlVecEnvWrapper(env, device=device, policy_obs_mode="actor") + torch_device = torch.device(device) + actor, hist_normalizer = build_student_actor_and_normalizer( + wrapped_env, + cfg, + device=torch_device, + ) + load_distilled_checkpoint(actor, hist_normalizer, load_path, device=torch_device) + actor.eval() + hist_normalizer.eval() + + if cfg.training.sim_backend == "motrix": + raise NotImplementedError( + "HORA distillation play_only currently supports offline MuJoCo video rendering only." + ) + + output_video = Path(load_path_dir) / "play_video_stage2.mp4" + print(f"Rendering video to {output_video}...") + print("Collecting physics states...") + with torch.inference_mode(): + render_play_mode( + env, + sim_backend=cfg.training.sim_backend, + render_spacing=float( + getattr(cfg.training, "render_spacing", getattr(env.cfg, "render_spacing", 1.0)) + ), + num_steps=int(cfg.training.play_steps), + output_video=output_video, + initialize=lambda: wrapped_env.reset()[0], + step=lambda obs: wrapped_env.step( + _student_policy(actor, hist_normalizer, obs, device=torch_device) + )[0], + camera_kwargs=_play_camera_kwargs(cfg), + ) + print("Done.") + return str(output_video) + + +@hydra.main(version_base="1.3", config_path="../conf/hora_distill", config_name="config") +def main(cfg: DictConfig) -> None: + ensure_registries() + + if cfg.training.device: + device = str(cfg.training.device) + elif torch.cuda.is_available(): + device = "cuda" + elif torch.backends.mps.is_available(): + device = "mps" + else: + device = "cpu" + + if cfg.training.play_only: + play_hora_distill(cfg, device) + return + + cfg = _apply_teacher_defaults(cfg) + teacher_algo_family, teacher_task = _get_teacher_owner_spec(cfg) + if teacher_algo_family is None or teacher_task is None: + raise ValueError("HORA distillation requires teacher.algo_family and teacher.task.") + + teacher_checkpoint, _ = _resolve_teacher_checkpoint_path(cfg) + if teacher_checkpoint is None: + raise FileNotFoundError( + "Could not resolve HORA teacher checkpoint. " + f"teacher.algo_family={teacher_algo_family!r} " + f"teacher.task={teacher_task!r}. " + "Set algo.load_run and optionally algo.checkpoint." + ) + + teacher_metadata = _teacher_run_metadata( + cfg, + teacher_algo_family=teacher_algo_family, + teacher_checkpoint=teacher_checkpoint, + ) + timestamp = datetime.datetime.now().strftime("%Y-%m-%d_%H-%M-%S") + log_root = Path(cfg.training.log_dir) if cfg.training.log_dir else get_log_root(ROOT_DIR, cfg) + run_name = f"{timestamp}_{cfg.training.sim_backend}_{teacher_metadata['run_slug']}" + log_dir = log_root / str(cfg.training.task_name) / run_name + logger = setup_logger(log_dir, "hora_distill", echo=str(cfg.training.logger) != "no_print") + _write_distill_run_config(log_dir, cfg=cfg, teacher_metadata=teacher_metadata) + logger.info( + "teacher_algo=%s teacher_task=%s teacher_checkpoint=%s", + teacher_metadata["algo_family"], + teacher_metadata["task"], + teacher_metadata["checkpoint_path"], + ) + + env = create_env( + cfg, + num_envs=int(cfg.algo.num_envs), + env_cfg_override=_build_env_cfg_override(cfg), + ) + wrapped_env = RslRlVecEnvWrapper(env, device=device, policy_obs_mode="actor") + trainer = HoraDistillationTrainer( + wrapped_env, + cfg, + device=device, + log_dir=log_dir, + teacher_checkpoint=teacher_checkpoint, + teacher_algo_family=teacher_algo_family, + teacher_metadata=teacher_metadata, + distill_runtime_cfg=_resolved_distill_runtime_cfg(cfg), + logger=logger, + ) + trainer.train() + + +if __name__ == "__main__": + main() diff --git a/scripts/train_rsl_rl.py b/scripts/train_rsl_rl.py index 8f33cc04f..b555d26de 100644 --- a/scripts/train_rsl_rl.py +++ b/scripts/train_rsl_rl.py @@ -16,6 +16,7 @@ if str(ROOT_DIR) not in sys.path: sys.path.insert(0, str(ROOT_DIR)) +from unilab.algos.torch.rsl_rl_runtime import resolve_rsl_rl_ppo_runtime from unilab.base.backend.xml import materialize_scene_visual_override from unilab.training import ( BackendAdapter, @@ -85,6 +86,22 @@ def _algo_config_dict(cfg: DictConfig) -> dict[str, Any]: return cast(dict[str, Any], train_cfg_raw) +def _resolve_ppo_wrapper_cls(rl_cfg: dict[str, Any]) -> type[RslRlVecEnvWrapper]: + """Resolve the VecEnv wrapper class from the owner-selected PPO runtime. + + Args: + rl_cfg: Resolved algorithm config dictionary from Hydra composition. + + Returns: + Wrapper class used to adapt the UniLab env contract to the active + RSL-RL PPO runtime. + """ + return resolve_rsl_rl_ppo_runtime( + rl_cfg, + default_wrapper_cls=RslRlVecEnvWrapper, + ).wrapper_cls + + def _format_play_checkpoint_error( cfg: DictConfig, *, @@ -123,6 +140,8 @@ def _format_play_checkpoint_error( def play_rsl_rl(cfg: DictConfig, device: str) -> str | None: """Play mode for RSL-RL.""" + rl_cfg = _algo_config_dict(cfg) + wrapper_cls = _resolve_ppo_wrapper_cls(rl_cfg) task_log_root = get_log_root(ROOT_DIR, cfg) / str(cfg.training.task_name) load_path, load_path_dir = parse_checkpoint_path(cfg, root_dir=ROOT_DIR) @@ -153,8 +172,8 @@ def play_rsl_rl(cfg: DictConfig, device: str) -> str | None: num_envs=cfg.training.play_env_num, env_cfg_override=env_cfg_override, ) - wrapped_env = RslRlVecEnvWrapper(env, device=device) - train_cfg = normalize_ppo_train_cfg(_algo_config_dict(cfg)) + wrapped_env = wrapper_cls(env, device=device) + train_cfg = normalize_ppo_train_cfg(rl_cfg) if "runner" not in train_cfg: train_cfg["runner"] = {} train_cfg["runner"]["logger"] = "none" @@ -273,6 +292,8 @@ def main(cfg: DictConfig) -> None: num_envs=cfg.algo.num_envs, env_cfg_override=env_cfg_override, ) + rl_cfg = _algo_config_dict(cfg) + wrapper_cls = _resolve_ppo_wrapper_cls(rl_cfg) nan_guard_cfg = getattr(cfg.training, "nan_guard", None) if nan_guard_cfg is not None and getattr(nan_guard_cfg, "enabled", False): @@ -290,9 +311,9 @@ def main(cfg: DictConfig) -> None: ) env.set_nan_guard(guard) - wrapped_env = RslRlVecEnvWrapper(env, device=device) + wrapped_env = wrapper_cls(env, device=device) - train_cfg = normalize_ppo_train_cfg(_algo_config_dict(cfg)) + train_cfg = normalize_ppo_train_cfg(rl_cfg) if "runner" not in train_cfg: train_cfg["runner"] = {} diff --git a/src/unilab/algos/torch/appo/runtime.py b/src/unilab/algos/torch/appo/runtime.py new file mode 100644 index 000000000..ad5b58bd8 --- /dev/null +++ b/src/unilab/algos/torch/appo/runtime.py @@ -0,0 +1,69 @@ +"""Runtime resolution helpers for APPO script assembly. + +This module keeps entrypoint scripts generic: they resolve an APPO runtime bundle +from owner config and then call the returned runner/play entrypoints without +knowing which concrete runtime implementation is active. +""" + +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass +from typing import Any + + +@dataclass(frozen=True) +class APPORuntime: + """Resolved APPO runtime entrypoints consumed by the generic script. + + Args: + runner_cls: Runner class used for APPO training mode. + play_fn: Play-mode callable used for checkpoint playback. + + Returns: + Immutable APPO runtime bundle selected from owner config. + """ + + runner_cls: type[Any] + play_fn: Callable[..., str | None] + + +def resolve_appo_runtime( + rl_cfg: dict[str, Any], + *, + default_play_fn: Callable[..., str | None], +) -> APPORuntime: + """Resolve the APPO runtime bundle from owner config. + + Args: + rl_cfg: Resolved algorithm config dictionary from Hydra composition. + default_play_fn: Generic APPO play function used when no custom runtime + resolver is selected by the owner config. + + Returns: + ``APPORuntime`` containing the train and play entrypoints for the + selected APPO runtime. + """ + runtime_resolver = rl_cfg.get("runtime_resolver") + if runtime_resolver in (None, ""): + from unilab.algos.torch.appo.runner import APPORunner + + return APPORuntime(runner_cls=APPORunner, play_fn=default_play_fn) + + from rsl_rl.utils import resolve_callable + + resolver = resolve_callable(str(runtime_resolver)) + runtime = resolver(rl_cfg) + if runtime is None: + raise ValueError( + f"APPO runtime resolver {runtime_resolver!r} returned None for rl_cfg runtime selection." + ) + + runner_cls = getattr(runtime, "runner_cls", None) + play_fn = getattr(runtime, "play_fn", None) + if runner_cls is None or play_fn is None: + raise TypeError( + f"APPO runtime resolver {runtime_resolver!r} must return an object with " + "'runner_cls' and 'play_fn' attributes." + ) + return APPORuntime(runner_cls=runner_cls, play_fn=play_fn) diff --git a/src/unilab/algos/torch/hora/__init__.py b/src/unilab/algos/torch/hora/__init__.py new file mode 100644 index 000000000..c78f952c2 --- /dev/null +++ b/src/unilab/algos/torch/hora/__init__.py @@ -0,0 +1,32 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +from .distill import HoraDistillationTrainer +from .models import HoraActorModel, HoraCriticModel, HoraSharedActorCritic +from .ppo import HoraPPO + +if TYPE_CHECKING: + from .appo import HoraAPPORunner, play_hora_appo + +__all__ = [ + "HoraActorModel", + "HoraAPPORunner", + "HoraCriticModel", + "HoraDistillationTrainer", + "HoraPPO", + "HoraSharedActorCritic", + "play_hora_appo", +] + + +def __getattr__(name: str) -> Any: + if name in {"HoraAPPORunner", "play_hora_appo"}: + from .appo import HoraAPPORunner, play_hora_appo + + exports = { + "HoraAPPORunner": HoraAPPORunner, + "play_hora_appo": play_hora_appo, + } + return exports[name] + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/src/unilab/algos/torch/hora/appo.py b/src/unilab/algos/torch/hora/appo.py new file mode 100644 index 000000000..d9beb2ba7 --- /dev/null +++ b/src/unilab/algos/torch/hora/appo.py @@ -0,0 +1,229 @@ +"""HORA-owned APPO entry helpers.""" + +from __future__ import annotations + +import os +from collections.abc import Callable +from copy import deepcopy +from dataclasses import dataclass +from typing import Any, cast + +import torch +from omegaconf import DictConfig + +from unilab.algos.torch.hora.appo_runner import HoraAPPORunner +from unilab.algos.torch.hora.rsl_rl_compat import ( + convert_config_v3_to_v4, + is_rsl_rl_v4, + is_rsl_rl_v5, +) +from unilab.base.observations import get_obs_dims +from unilab.training import BackendAdapter, create_env +from unilab.visualization import render_play_mode + +from .observations import build_hora_actor_tensordict, split_hora_obs_with_priv_info +from .runtime import is_hora_appo_runtime + + +@dataclass(frozen=True) +class HoraAPPORuntime: + """Resolved HORA APPO entrypoints used by the generic APPO script. + + Args: + runner_cls: Runner class used for HORA APPO training mode. + play_fn: Play-mode callable used for HORA APPO checkpoint playback. + + Returns: + Immutable entrypoint bundle consumed by generic APPO script assembly. + """ + + runner_cls: type[HoraAPPORunner] + play_fn: Callable[..., str | None] + + +def resolve_hora_appo_runtime(rl_cfg: dict[str, Any]) -> HoraAPPORuntime | None: + """Resolve HORA APPO entrypoints from an explicit runtime marker. + + Args: + rl_cfg: Resolved algorithm config dictionary from Hydra composition. + + Returns: + ``HoraAPPORuntime`` when the owner config selects HORA APPO, otherwise + ``None``. + """ + if not is_hora_appo_runtime(rl_cfg): + return None + return HoraAPPORuntime(runner_cls=HoraAPPORunner, play_fn=play_hora_appo) + + +def _update_hora_obs_groups( + rl_cfg: dict[str, Any], + *, + obs_dim: int, + priv_info_dim: int, +) -> None: + """Update grouped actor/critic dims for the HORA APPO runtime. + + Args: + rl_cfg: Mutable algorithm config dictionary to update in place. + obs_dim: Actor observation dimension reported by the env contract. + priv_info_dim: Privileged-info dimension reported by the env contract. + + Returns: + None. Mutates ``rl_cfg["obs_groups"]`` directly. + """ + obs_groups = rl_cfg.setdefault("obs_groups", {}) + actor_group = obs_groups.setdefault("actor", {}) + critic_group = obs_groups.setdefault("critic", {}) + if isinstance(actor_group, dict): + actor_group["actor"] = obs_dim + actor_group["priv_info"] = priv_info_dim + if isinstance(critic_group, dict): + critic_group["actor"] = obs_dim + critic_group["priv_info"] = priv_info_dim + + +def play_hora_appo( + cfg: DictConfig, + rl_cfg: dict[str, Any], + *, + root_dir, + resolve_checkpoint_path, +) -> str | None: + """Play HORA APPO checkpoints with grouped actor and privileged inputs.""" + import numpy as np + from rsl_rl.utils import resolve_callable + from tensordict import TensorDict + + env_cfg_override = BackendAdapter( + cfg, + root_dir=root_dir, + algo_name="appo", + ).build_task_env_cfg_override() + + device = cfg.training.device or ( + "cuda" + if torch.cuda.is_available() + else "mps" + if torch.backends.mps.is_available() + else "cpu" + ) + print(f"Using device for play: {device}") + + env = cast( + Any, + create_env( + cfg, + num_envs=cfg.training.play_env_num, + env_cfg_override=env_cfg_override, + ), + ) + obs_dim, _ = get_obs_dims(env.obs_groups_spec) + if env.state is None: + env.init_state() + _, _, state_priv_info = split_hora_obs_with_priv_info( + env.state.obs, + env.state.info if env.state is not None else None, + ) + priv_info_dim = int(state_priv_info.shape[1]) if state_priv_info is not None else 0 + if priv_info_dim <= 0: + raise ValueError("HORA APPO play requires privileged info from the environment.") + + action_shape = env.action_space.shape + if action_shape is None: + raise ValueError("env.action_space.shape must be defined") + action_dim = int(action_shape[0]) + + rl_cfg_dict = dict(rl_cfg) + _update_hora_obs_groups(rl_cfg_dict, obs_dim=obs_dim, priv_info_dim=priv_info_dim) + + if is_rsl_rl_v5(): + pass + elif is_rsl_rl_v4(): + rl_cfg_dict = convert_config_v3_to_v4(rl_cfg_dict) + + obs_example = torch.zeros((cfg.training.play_env_num, obs_dim), device=device) + td_example = TensorDict( + { + "actor": obs_example, + "priv_info": torch.zeros((cfg.training.play_env_num, priv_info_dim), device=device), + }, + batch_size=cfg.training.play_env_num, + ) + + actor_cfg = deepcopy(rl_cfg_dict["actor"]) + actor_cls = resolve_callable(actor_cfg.pop("class_name")) + actor_cfg.pop("num_actions", None) + actor = actor_cls(td_example, rl_cfg_dict["obs_groups"], "actor", action_dim, **actor_cfg) + actor = actor.to(device) + actor.eval() + + load_path, load_path_dir = resolve_checkpoint_path(cfg) + if not load_path or not os.path.exists(load_path): + print(f"Could not find run to load. load_path={load_path}") + return None + + print(f"Loading model: {load_path}") + checkpoint = torch.load(load_path, map_location=device, weights_only=True) + actor.load_state_dict(checkpoint["actor"]) + + if load_path_dir is None: + print(f"Could not resolve checkpoint directory. load_path_dir={load_path_dir}") + return None + + output_video = os.path.join(load_path_dir, "play_video.mp4") + print(f"Rendering video to {output_video}...") + num_steps = int(getattr(cfg.training, "play_steps", 1000)) + current_priv_info: np.ndarray | None = None + + def initialize_play_obs() -> np.ndarray: + nonlocal current_priv_info + obs_out, info_out = env.reset(np.arange(cfg.training.play_env_num, dtype=np.int32)) + actor_obs, _, priv_info = split_hora_obs_with_priv_info(obs_out, info_out) + current_priv_info = priv_info.astype(np.float32) if priv_info is not None else None + return np.asarray(actor_obs, dtype=np.float32) + + def step_play_obs(obs_np: np.ndarray) -> np.ndarray: + nonlocal current_priv_info + if current_priv_info is None: + raise ValueError("HORA APPO play step is missing privileged info.") + td = build_hora_actor_tensordict( + obs_np, + priv_info=current_priv_info, + device=device, + batch_size=cfg.training.play_env_num, + ) + actions = actor(td).cpu().numpy().astype(np.float32) + state = env.step(actions) + actor_obs, _, priv_info = split_hora_obs_with_priv_info(state.obs, state.info) + current_priv_info = priv_info.astype(np.float32) if priv_info is not None else None + return np.asarray(actor_obs, dtype=np.float32) + + print("Collecting physics states...") + with torch.inference_mode(): + render_play_mode( + env, + sim_backend=cfg.training.sim_backend, + render_spacing=float( + getattr(cfg.training, "render_spacing", getattr(env.cfg, "render_spacing", 1.0)) + ), + num_steps=num_steps, + output_video=output_video, + initialize=initialize_play_obs, + step=step_play_obs, + camera_kwargs={ + "cam_distance": cfg.training.cam_distance, + "cam_elevation": cfg.training.cam_elevation, + "cam_azimuth": cfg.training.cam_azimuth, + "cam_lookat": getattr(cfg.training, "cam_lookat", None), + "cam_tracking": getattr(cfg.training, "cam_tracking", False), + "cam_tracking_env_idx": getattr(cfg.training, "cam_tracking_env_idx", 0), + "cam_tracking_extra_envs": getattr(cfg.training, "cam_tracking_extra_envs", 2), + }, + ) + print(f"Saving video to {output_video} with mediapy...") + print("Done.") + return output_video + + +__all__ = ["HoraAPPORunner", "HoraAPPORuntime", "play_hora_appo", "resolve_hora_appo_runtime"] diff --git a/src/unilab/algos/torch/hora/appo_learner.py b/src/unilab/algos/torch/hora/appo_learner.py new file mode 100644 index 000000000..dd30758fa --- /dev/null +++ b/src/unilab/algos/torch/hora/appo_learner.py @@ -0,0 +1,270 @@ +"""HORA-owned APPO learner with grouped actor and privileged observations.""" + +from __future__ import annotations + +from itertools import chain + +import torch +import torch.nn as nn +from tensordict import TensorDict + +from unilab.algos.torch.appo.learner import APPOLearner, vtrace_advantages + + +def _build_hora_obs_td( + actor_obs: torch.Tensor, + *, + device: str, + priv_info: torch.Tensor | None = None, +) -> TensorDict: + if priv_info is not None: + return TensorDict( + {"actor": actor_obs, "priv_info": priv_info}, + batch_size=actor_obs.shape[0], + device=device, + ) + return TensorDict({"policy": actor_obs}, batch_size=actor_obs.shape[0], device=device) + + +def _derive_priv_info_from_critic( + actor_obs: torch.Tensor, + critic_obs: torch.Tensor | None, + *, + context: str, +) -> torch.Tensor: + if critic_obs is None: + raise ValueError(f"HORA APPO {context} requires critic observations.") + actor_dim = int(actor_obs.shape[-1]) + critic_dim = int(critic_obs.shape[-1]) + if critic_dim <= actor_dim: + raise ValueError( + f"HORA APPO {context} requires critic observations to include privileged tail " + f"features; got actor_dim={actor_dim}, critic_dim={critic_dim}." + ) + return critic_obs[..., actor_dim:] + + +class HoraAPPOLearner(APPOLearner): + """APPO learner variant for HORA grouped observations.""" + + def process_batch(self, batch_dict): + """Compute V-trace targets for grouped HORA rollouts.""" + obs = batch_dict["observations"] + critic_base = batch_dict.get("critic", None) + rewards = batch_dict["rewards"] + dones = batch_dict["dones"].float() + last_obs = batch_dict["last_obs"] + last_critic = batch_dict.get("last_critic", None) + behavior_log_probs = batch_dict["actions_log_prob"] + actions = batch_dict["actions"] + + T, N = obs.shape[:2] + priv_info = _derive_priv_info_from_critic( + obs, + critic_base, + context="rollout batch", + ) + last_priv_info = _derive_priv_info_from_critic( + last_obs, + last_critic, + context="bootstrap batch", + ) + obs_flat = obs.flatten(0, 1) + priv_info_flat = priv_info.flatten(0, 1) + + obs_td = _build_hora_obs_td(obs_flat, device=self.device, priv_info=priv_info_flat) + last_obs_td = _build_hora_obs_td(last_obs, device=self.device, priv_info=last_priv_info) + + critic_obs = critic_base + critic_obs_flat = critic_obs.flatten(0, 1) + critic_obs_td = _build_hora_obs_td(obs_flat, device=self.device, priv_info=priv_info_flat) + critic_last_obs_td = _build_hora_obs_td( + last_obs, + device=self.device, + priv_info=last_priv_info, + ) + + if hasattr(self.actor, "update_normalization"): + self.actor.update_normalization(obs_td) + self.actor.update_normalization(last_obs_td) + if hasattr(self.critic, "update_normalization"): + self.critic.update_normalization(critic_obs_td) + self.critic.update_normalization(critic_last_obs_td) + + batch_dict["_critic_obs_flat"] = critic_obs_flat + batch_dict["_critic_obs_td"] = critic_obs_td + + with torch.inference_mode(): + values_flat = self.critic(critic_obs_td) + last_values = self.critic(critic_last_obs_td).squeeze(-1) + values = values_flat.view(T, N, -1).squeeze(-1) + + actions_flat = actions.flatten(0, 1) + with torch.inference_mode(): + self.target_actor(obs_td, stochastic_output=True) + target_log_probs_flat = self.target_actor.get_output_log_prob(actions_flat) + batch_dict["_old_mu"] = self.target_actor.output_mean.clone() + batch_dict["_old_sigma"] = self.target_actor.output_std.clone() + target_log_probs = target_log_probs_flat.view(T, N) + + vs, advantages = vtrace_advantages( + behavior_log_probs=behavior_log_probs, + target_log_probs=target_log_probs, + rewards=rewards, + values=values, + bootstrap_values=last_values, + dones=dones, + gamma=self.gamma, + clip_rho=self.vtrace_clip_rho, + clip_c=self.vtrace_clip_c, + ) + + batch_dict["values"] = values + batch_dict["advantages"] = advantages + batch_dict["returns"] = vs + batch_dict["target_log_probs"] = target_log_probs + batch_dict["_obs_td"] = obs_td + + return batch_dict + + def update(self, batch_dict): + """Perform APPO update for grouped HORA observations.""" + obs_flat = batch_dict["observations"].flatten(0, 1) + priv_info = _derive_priv_info_from_critic( + batch_dict["observations"], + batch_dict.get("critic"), + context="update batch", + ) + priv_info_flat = priv_info.flatten(0, 1) + actions_flat = batch_dict["actions"].flatten(0, 1) + returns_flat = batch_dict["returns"].flatten(0, 1) + advantages_flat = batch_dict["advantages"].flatten(0, 1) + behavior_log_probs_flat = batch_dict["actions_log_prob"].flatten(0, 1) + old_values_flat = batch_dict["values"].flatten(0, 1) + target_log_probs_flat = batch_dict["target_log_probs"].flatten(0, 1) + advantages_flat = (advantages_flat - advantages_flat.mean()) / ( + advantages_flat.std() + 1e-8 + ) + + obs_td = batch_dict.get("_obs_td") + if obs_td is None: + obs_td = _build_hora_obs_td(obs_flat, device=self.device, priv_info=priv_info_flat) + + critic_obs_td = batch_dict.get("_critic_obs_td") + if critic_obs_td is None: + critic_obs_td = _build_hora_obs_td( + obs_flat, + device=self.device, + priv_info=priv_info_flat, + ) + + with torch.inference_mode(): + old_mu_flat = batch_dict["_old_mu"] + old_sigma_flat = batch_dict["_old_sigma"] + + batch_size = obs_flat.shape[0] + mini_batch_size = batch_size // self.num_mini_batches + + mean_value_loss = 0.0 + mean_surrogate_loss = 0.0 + mean_entropy = 0.0 + mean_kl = 0.0 + num_updates = 0 + + for _epoch in range(self.num_learning_epochs): + indices = torch.randperm(batch_size, device=self.device) + + for i in range(self.num_mini_batches): + start = i * mini_batch_size + end = (i + 1) * mini_batch_size + batch_idx = indices[start:end] + + obs_mini_td = obs_td[batch_idx] + critic_obs_mini_td = critic_obs_td[batch_idx] + actions_mini = actions_flat[batch_idx] + target_values_mini = returns_flat[batch_idx] + advantages_mini = advantages_flat[batch_idx] + behavior_logp_mini = behavior_log_probs_flat[batch_idx] + old_values_mini = old_values_flat[batch_idx] + target_logp_mini = target_log_probs_flat[batch_idx] + old_mu_mini = old_mu_flat[batch_idx] + old_sigma_mini = old_sigma_flat[batch_idx] + + _ = self.actor(obs_mini_td, stochastic_output=True) + current_log_prob = self.actor.get_output_log_prob(actions_mini) + value = self.critic(critic_obs_mini_td).squeeze(-1) + entropy = self.actor.output_entropy.mean() + + mu = self.actor.output_mean + sigma = self.actor.output_std + + with torch.no_grad(): + clipped_rho = torch.clamp( + torch.exp(behavior_logp_mini - target_logp_mini), max=1.0 + ) + ratio = clipped_rho * torch.exp(current_log_prob - behavior_logp_mini) + + surrogate = -advantages_mini * ratio + surrogate_clipped = -advantages_mini * torch.clamp( + ratio, 1.0 - self.clip_param, 1.0 + self.clip_param + ) + surrogate_loss = torch.max(surrogate, surrogate_clipped).mean() + + if self.desired_kl is not None and self.schedule == "adaptive": + with torch.inference_mode(): + kl = torch.sum( + torch.log(sigma / old_sigma_mini + 1e-5) + + (old_sigma_mini.pow(2) + (old_mu_mini - mu).pow(2)) + / (2.0 * sigma.pow(2)) + - 0.5, + dim=-1, + ) + kl_mean = torch.mean(kl) + + if kl_mean > self.desired_kl * 2.0: + self.learning_rate = max(1e-5, self.learning_rate / 1.5) + elif kl_mean < self.desired_kl / 2.0 and kl_mean > 0.0: + self.learning_rate = min(1e-2, self.learning_rate * 1.5) + + for param_group in self.optimizer.param_groups: + param_group["lr"] = self.learning_rate + + mean_kl += kl_mean.item() + + if self.use_clipped_value_loss: + value_clipped = old_values_mini + (value - old_values_mini).clamp( + -self.clip_param, self.clip_param + ) + value_losses = (value - target_values_mini).pow(2) + value_losses_clipped = (value_clipped - target_values_mini).pow(2) + value_loss = torch.max(value_losses, value_losses_clipped).mean() + else: + value_loss = (value - target_values_mini).pow(2).mean() + + loss = ( + surrogate_loss + self.value_loss_coef * value_loss - self.entropy_coef * entropy + ) + + self.optimizer.zero_grad(set_to_none=True) + loss.backward() + nn.utils.clip_grad_norm_( + chain(self.actor.parameters(), self.critic.parameters()), self.max_grad_norm + ) + self.optimizer.step() + + mean_value_loss += value_loss.item() + mean_surrogate_loss += surrogate_loss.item() + mean_entropy += entropy.item() + num_updates += 1 + + self._update_counter += 1 + if self._update_counter % self.target_update_freq == 0: + self.update_target_network() + + num_updates = max(num_updates, 1) + return { + "surrogate_loss": mean_surrogate_loss / num_updates, + "value_loss": mean_value_loss / num_updates, + "entropy": mean_entropy / num_updates, + "kl": mean_kl / num_updates if self.schedule == "adaptive" else 0.0, + } diff --git a/src/unilab/algos/torch/hora/appo_runner.py b/src/unilab/algos/torch/hora/appo_runner.py new file mode 100644 index 000000000..3d4d67282 --- /dev/null +++ b/src/unilab/algos/torch/hora/appo_runner.py @@ -0,0 +1,342 @@ +"""HORA-owned APPO runner.""" + +from __future__ import annotations + +import multiprocessing as mp +import os +import time +from collections import deque +from copy import deepcopy +from typing import Any + +import numpy as np +import torch +from rsl_rl.utils import resolve_callable + +from unilab.algos.torch.appo.runner import APPORunner +from unilab.algos.torch.hora.appo_learner import HoraAPPOLearner +from unilab.algos.torch.hora.appo_worker import hora_appo_collector_fn +from unilab.algos.torch.hora.rsl_rl_compat import ( + convert_config_v3_to_v4, + is_rsl_rl_v4, + is_rsl_rl_v5, +) +from unilab.base.observations import get_critic_base_dim, get_obs_dims +from unilab.base.registry import ensure_registries +from unilab.ipc import SharedOnPolicyStorage, SharedWeightSync +from unilab.logging import OffPolicyLogger + + +class HoraAPPORunner(APPORunner): + """APPO runner variant that preserves grouped HORA observations.""" + + def __init__(self, *args, **kwargs): + self.priv_info_dim = 0 + super().__init__(*args, **kwargs) + + def _resolve_dims(self): + self.obs_dim, self.action_dim = self._detect_dims() + + obs_groups = self.rl_cfg.setdefault("obs_groups", {}) + actor_group = obs_groups.setdefault("actor", {}) + critic_group = obs_groups.setdefault("critic", {}) + if isinstance(actor_group, dict): + actor_group["actor"] = self.obs_dim + actor_group["priv_info"] = self.priv_info_dim + if isinstance(critic_group, dict): + critic_group["actor"] = self.obs_dim + critic_group["priv_info"] = self.priv_info_dim + + def _detect_dims(self): + from unilab.base import registry + + ensure_registries() + + env = registry.make( + self.env_name, + num_envs=self._detect_dim_probe_num_envs(), + sim_backend=self.sim_backend, + env_cfg_override=self.env_cfg_overrides if self.env_cfg_overrides else None, + ) + obs_dim, critic_dim = get_obs_dims(env.obs_groups_spec) + self.critic_dim = critic_dim + self.critic_input_dim = get_critic_base_dim(env.obs_groups_spec) + if env.state is None: + env.init_state() + info = env.state.info if env.state is not None else {} + priv_info = info.get("critic_info") if isinstance(info, dict) else None + if isinstance(priv_info, np.ndarray) and priv_info.ndim == 2: + self.priv_info_dim = int(priv_info.shape[1]) + elif critic_dim > obs_dim: + self.priv_info_dim = int(critic_dim - obs_dim) + if self.priv_info_dim <= 0: + env.close() + raise ValueError("HORA APPO requires a positive privileged-info dimension.") + assert env.action_space.shape is not None + action_dim = env.action_space.shape[0] + env.close() + return obs_dim, action_dim + + def _detect_dim_probe_num_envs(self) -> int: + scale_list = None + if isinstance(self.env_cfg_overrides, dict): + domain_rand = self.env_cfg_overrides.get("domain_rand") + if isinstance(domain_rand, dict): + scale_list = domain_rand.get("scale_list") + if scale_list is None: + scale_list = self.env_cfg_overrides.get("scale_list") + if isinstance(scale_list, (list, tuple)): + return max(1, len(scale_list)) + return 1 + + def _build_learner(self): + cfg = dict(self.rl_cfg) + if is_rsl_rl_v5(): + pass + elif is_rsl_rl_v4(): + cfg = convert_config_v3_to_v4(cfg) + + from tensordict import TensorDict + + obs_example = torch.zeros((self.num_envs, self.obs_dim), device=self.device) + priv_info_example = torch.zeros((self.num_envs, self.priv_info_dim), device=self.device) + td_example = TensorDict( + {"actor": obs_example, "priv_info": priv_info_example}, + batch_size=self.num_envs, + device=self.device, + ) + + actor_cfg = deepcopy(cfg.get("actor", {})) + actor_cls = resolve_callable(actor_cfg.pop("class_name")) + actor_cfg.pop("num_actions", None) + actor = actor_cls(td_example, cfg["obs_groups"], "actor", self.action_dim, **actor_cfg) + + critic_cfg: dict[str, Any] = deepcopy(cfg.get("critic") or cfg.get("actor") or {}) + critic_cls = resolve_callable(critic_cfg.pop("class_name", "rsl_rl.models.MLPModel")) + critic_cfg.pop("num_actions", None) + critic_cfg.pop("distribution_cfg", None) + critic = critic_cls(td_example, cfg["obs_groups"], "critic", 1, **critic_cfg) + + algo_cfg = cfg.get("algorithm", cfg) + return HoraAPPOLearner( + actor=actor, + critic=critic, + device=self.device, + num_learning_epochs=algo_cfg.get("num_learning_epochs", 5), + num_mini_batches=algo_cfg.get("num_mini_batches", 4), + clip_param=algo_cfg.get("clip_param", 0.2), + gamma=algo_cfg.get("gamma", 0.99), + lam=algo_cfg.get("lam", 0.95), + value_loss_coef=algo_cfg.get("value_loss_coef", 1.0), + entropy_coef=algo_cfg.get("entropy_coef", 0.01), + learning_rate=algo_cfg.get("learning_rate", 1e-3), + max_grad_norm=algo_cfg.get("max_grad_norm", 1.0), + use_clipped_value_loss=algo_cfg.get("use_clipped_value_loss", True), + schedule=algo_cfg.get("schedule", "fixed"), + desired_kl=algo_cfg.get("desired_kl", 0.01), + optimizer=algo_cfg.get("optimizer", "adam"), + tau=algo_cfg.get("tau", 1.0), + target_update_freq=algo_cfg.get("target_update_freq", 1), + vtrace_clip_rho=algo_cfg.get("vtrace_clip_rho", 1.0), + vtrace_clip_c=algo_cfg.get("vtrace_clip_c", 1.0), + ) + + def _collector_fn(self, stop_event, **kwargs): + hora_appo_collector_fn(stop_event=stop_event, **kwargs) + + def learn( + self, + max_iterations: int = 1500, + save_interval: int = 50, + log_dir: str = "logs", + logger_type: str = "tensorboard", + ) -> None: + os.makedirs(log_dir, exist_ok=True) + train_start_wall = time.time() + best_mean_reward = float("-inf") + last_mean_reward = 0.0 + ckpt_path: str | None = None + iteration = 0 + + learner = self._build_learner() + + shared_storage = SharedOnPolicyStorage( + num_envs=self.num_envs, + num_steps=self.steps_per_env, + obs_dim=self.obs_dim, + action_dim=self.action_dim, + critic_dim=self.critic_dim, + num_slots=4, + create=True, + ) + self._shared_resources.append(shared_storage) + + actor_weight_sync = SharedWeightSync.from_state_dict( + learner.actor.state_dict(), create=True + ) + critic_weight_sync = SharedWeightSync.from_state_dict( + learner.critic.state_dict(), + create=True, + ) + self._shared_resources.extend([actor_weight_sync, critic_weight_sync]) + + actor_weight_param_shapes = { + name: p.shape for name, p in learner.actor.state_dict().items() + } + critic_weight_param_shapes = { + name: p.shape for name, p in learner.critic.state_dict().items() + } + + metrics_queue: mp.Queue = mp.get_context("spawn").Queue(maxsize=100) + collector_kwargs = { + "env_name": self.env_name, + "rl_cfg": self.rl_cfg, + "num_envs": self.num_envs, + "steps_per_env": self.steps_per_env, + "shm_storage_name": shared_storage.name, + "sync_primitives": ( + shared_storage._write_ptr, + shared_storage._read_ptr, + ), + "obs_dim": self.obs_dim, + "action_dim": self.action_dim, + "critic_dim": self.critic_dim, + "priv_info_dim": self.priv_info_dim, + "actor_weight_sync_name": actor_weight_sync.name, + "actor_weight_param_shapes": actor_weight_param_shapes, + "critic_weight_sync_name": critic_weight_sync.name, + "critic_weight_param_shapes": critic_weight_param_shapes, + "metrics_queue": metrics_queue, + "collector_device": self.collector_device, + "sim_backend": self.sim_backend, + "env_cfg_override": self.env_cfg_overrides if self.env_cfg_overrides else None, + } + self._start_collector( + target_fn=hora_appo_collector_fn, + kwargs={"stop_event": self._stop_event, **collector_kwargs}, + ) + + env_steps_per_sync = self.steps_per_env * self.num_envs + logger = OffPolicyLogger( + algo_name="APPO", + max_iterations=max_iterations, + num_envs=self.num_envs, + env_name=self.env_name, + obs_dim=self.obs_dim, + action_dim=self.action_dim, + log_dir=log_dir, + log_backend=logger_type, + ) + logger.set_collection_sync(True, env_steps_per_sync) + logger.start() + logger.log_status( + f"Waiting for first rollout... " + f"(replay_queue={self.replay_queue_size}, " + f"epochs={learner.num_learning_epochs})" + ) + + reward_history: deque = deque(maxlen=200) + latest_reward_components: dict = {} + replay_queue: deque[dict] = deque(maxlen=self.replay_queue_size) + + for iteration in range(1, max_iterations + 1): + iter_start = time.time() + self._drain_metrics(metrics_queue, reward_history, latest_reward_components, logger) + wait_start = time.time() + + data_ready = shared_storage.wait_for_data(timeout=60.0) + if not data_ready: + if not self._check_collector_alive(): + self._drain_metrics( + metrics_queue, + reward_history, + latest_reward_components, + logger, + ) + raise RuntimeError( + "APPO collector process died before producing data. " + "Check stderr for [HORA APPO WORKER CRASH] messages." + ) + logger.log_status( + f"[yellow]Warning: Timeout waiting for data at iteration {iteration}[/]" + ) + continue + + available_on_arrive = shared_storage.available() + wait_time = time.time() - wait_start + + num_new = shared_storage.available() + for _ in range(num_new): + raw = shared_storage.read_torch(self.device) + shared_storage.advance_read() + + rollout: dict = {} + for k, v in raw.items(): + if k not in ("last_obs", "last_critic") and v.ndim >= 2: + rollout[k] = v.transpose(0, 1) + else: + rollout[k] = v + if "obs" in rollout: + rollout["observations"] = rollout.pop("obs") + if "log_probs" in rollout: + rollout["actions_log_prob"] = rollout.pop("log_probs") + replay_queue.append(rollout) + + self._drain_metrics(metrics_queue, reward_history, latest_reward_components, logger) + collect_time = time.time() - iter_start + + combined: dict = {} + for k in replay_queue[0]: + if k in ("last_obs", "last_critic"): + combined[k] = torch.cat([r[k] for r in replay_queue], dim=0) + else: + combined[k] = torch.cat([r[k] for r in replay_queue], dim=1) + + train_start = time.time() + learner.process_batch(combined) + metrics = learner.update(combined) + actor_weight_sync.write_weights(learner.actor.state_dict()) + critic_weight_sync.write_weights(learner.critic.state_dict()) + train_time = time.time() - train_start + + metrics["replay_queue_len"] = float(len(replay_queue)) + metrics["available_on_arrive"] = float(available_on_arrive) + logger.update_replay_queue(len(replay_queue), self.replay_queue_size) + + mean_reward = ( + sum(list(reward_history)[-50:]) / max(len(list(reward_history)[-50:]), 1) + if reward_history + else 0.0 + ) + last_mean_reward = float(mean_reward) + best_mean_reward = max(best_mean_reward, last_mean_reward) + + logger.log_step( + iteration=iteration, + metrics=metrics, + reward=mean_reward, + reward_components=latest_reward_components, + collect_time=collect_time, + train_time=train_time, + wait_time=wait_time, + ) + + if save_interval > 0 and iteration % save_interval == 0: + ckpt_path = os.path.join(log_dir, f"model_{iteration}.pt") + torch.save(learner.get_state_dict(), ckpt_path) + logger.log_save(ckpt_path) + + ckpt_path = os.path.join(log_dir, f"model_{max_iterations}.pt") + torch.save(learner.get_state_dict(), ckpt_path) + logger.log_save(ckpt_path) + logger.finish() + self.last_run_summary = { + "status": "completed", + "completed_iterations": iteration, + "total_env_steps": int(logger._total_steps), + "final_mean_reward": last_mean_reward if reward_history else None, + "best_mean_reward": best_mean_reward if reward_history else None, + "mean_episode_length": float(logger._mean_ep_length), + "last_checkpoint": ckpt_path, + "training_wall_time_sec": time.time() - train_start_wall, + } diff --git a/src/unilab/algos/torch/hora/appo_worker.py b/src/unilab/algos/torch/hora/appo_worker.py new file mode 100644 index 000000000..a3b393a4d --- /dev/null +++ b/src/unilab/algos/torch/hora/appo_worker.py @@ -0,0 +1,382 @@ +"""HORA-owned APPO rollout worker.""" + +from __future__ import annotations + +import statistics +import sys +import time +from collections import defaultdict +from typing import Any, Dict + +import numpy as np +import torch +from rsl_rl.utils import resolve_callable + +from unilab.base.final_observation import resolve_terminal_observation_contract +from unilab.base.registry import ensure_registries + +from .observations import split_hora_obs_with_priv_info + + +def compute_hora_timeout_bootstrap_correction( + critic: Any, + collector_device: str, + gamma: float, + timeout_mask: np.ndarray, + final_obs: np.ndarray, + final_critic: np.ndarray | None = None, + final_priv_info: np.ndarray | None = None, +) -> np.ndarray: + """Compute timeout bootstrap values for grouped HORA observations.""" + corrections = np.zeros(timeout_mask.shape, dtype=np.float32) + if not np.any(timeout_mask): + return corrections + + from tensordict import TensorDict + + if final_priv_info is not None: + actor_input = torch.from_numpy(final_obs[timeout_mask]).to(collector_device) + priv_input = torch.from_numpy(final_priv_info[timeout_mask]).to(collector_device) + critic_td = TensorDict( + {"actor": actor_input, "priv_info": priv_input}, + batch_size=actor_input.shape[0], + device=collector_device, + ) + else: + critic_input_np = final_critic if final_critic is not None else final_obs + critic_input = torch.from_numpy(critic_input_np[timeout_mask]).to(collector_device) + critic_td = TensorDict( + {"policy": critic_input}, + batch_size=critic_input.shape[0], + device=collector_device, + ) + with torch.no_grad(): + bootstrap = critic(critic_td).squeeze(-1).cpu().numpy().astype(np.float32, copy=False) + corrections[timeout_mask] = float(gamma) * bootstrap + return corrections + + +def hora_appo_collector_fn( + stop_event: Any, + env_name: str, + rl_cfg: dict, + num_envs: int, + steps_per_env: int, + shm_storage_name: Dict[str, str], + sync_primitives: tuple, + obs_dim: int, + action_dim: int, + critic_dim: int, + actor_weight_sync_name: str, + actor_weight_param_shapes: dict, + critic_weight_sync_name: str, + critic_weight_param_shapes: dict, + metrics_queue: Any, + collector_device: str = "cpu", + sim_backend: str = "mujoco", + env_cfg_override: dict | None = None, + priv_info_dim: int = 0, +): + """Collect grouped HORA APPO rollouts into shared storage.""" + from copy import deepcopy + + from tensordict import TensorDict + + from unilab.algos.torch.hora.rsl_rl_compat import ( + convert_config_v3_to_v4, + is_rsl_rl_v4, + is_rsl_rl_v5, + ) + from unilab.base import registry + from unilab.ipc import SharedOnPolicyStorage, SharedWeightSync + + ensure_registries() + + storage = SharedOnPolicyStorage( + num_envs=num_envs, + num_steps=steps_per_env, + obs_dim=obs_dim, + action_dim=action_dim, + critic_dim=critic_dim, + create=False, + shm_name_prefix=shm_storage_name, + ) + storage.attach_sync_primitives(*sync_primitives) + actor_weight_sync = SharedWeightSync( + actor_weight_param_shapes, + create=False, + shm_name=actor_weight_sync_name, + ) + critic_weight_sync = SharedWeightSync( + critic_weight_param_shapes, + create=False, + shm_name=critic_weight_sync_name, + ) + + env: Any = registry.make( + env_name, + num_envs=num_envs, + sim_backend=sim_backend, + env_cfg_override=env_cfg_override, + ) + + cfg = dict(rl_cfg) + if is_rsl_rl_v5(): + pass + elif is_rsl_rl_v4(): + cfg = convert_config_v3_to_v4(cfg) + + if priv_info_dim <= 0: + raise ValueError("HORA APPO collector requires priv_info_dim > 0") + + obs_example = torch.zeros((num_envs, obs_dim), device=collector_device) + priv_info_example = torch.zeros((num_envs, priv_info_dim), device=collector_device) + td_example = TensorDict( + {"actor": obs_example, "priv_info": priv_info_example}, + batch_size=num_envs, + device=collector_device, + ) + + actor_cfg = deepcopy(cfg["actor"]) + actor_cls = resolve_callable(actor_cfg.pop("class_name")) + actor_cfg.pop("num_actions", None) + actor = actor_cls(td_example, cfg["obs_groups"], "actor", action_dim, **actor_cfg) + actor = actor.to(collector_device) + actor.eval() + + critic_cfg = deepcopy(cfg.get("critic") or cfg.get("actor") or {}) + critic_cls = resolve_callable(critic_cfg.pop("class_name", "rsl_rl.models.MLPModel")) + critic_cfg.pop("num_actions", None) + critic_cfg.pop("distribution_cfg", None) + critic = critic_cls(td_example, cfg["obs_groups"], "critic", 1, **critic_cfg) + critic = critic.to(collector_device) + critic.eval() + + actor_sd = dict(actor.state_dict()) + actor_weight_sync.read_weights_into(actor_sd) + actor.load_state_dict(actor_sd) + local_actor_weight_version = actor_weight_sync.version + + critic_sd = dict(critic.state_dict()) + critic_weight_sync.read_weights_into(critic_sd) + critic.load_state_dict(critic_sd) + local_critic_weight_version = critic_weight_sync.version + + env_indices = np.arange(num_envs, dtype=np.int32) + try: + obs_out, info_out = env.reset(env_indices) + except TypeError: + obs_out, info_out = env.reset() + + def to_float32_np(x): + if hasattr(x, "cpu"): + x = x.cpu().numpy() + return np.asarray(x, dtype=np.float32) + + obs_np, critic_np, priv_info_np = split_hora_obs_with_priv_info(obs_out, info_out) + obs_np = to_float32_np(obs_np) + if critic_np is not None: + critic_np = to_float32_np(critic_np) + if priv_info_np is not None: + priv_info_np = to_float32_np(priv_info_np) + if priv_info_np is None: + raise ValueError("HORA APPO collector did not receive privileged info from env reset.") + + obs_torch = torch.zeros((num_envs, obs_dim), dtype=torch.float32, device=collector_device) + priv_info_torch = torch.zeros( + (num_envs, priv_info_dim), + dtype=torch.float32, + device=collector_device, + ) + obs_td = TensorDict( + {"actor": obs_torch, "priv_info": priv_info_torch}, + batch_size=num_envs, + device=collector_device, + ) + + total_steps = 0 + ep_rewards = [] + ep_lengths = [] + current_ep_rewards = np.zeros(num_envs, dtype=np.float32) + current_ep_lengths = np.zeros(num_envs, dtype=np.int32) + ep_reward_components = defaultdict(list) + ep_timeouts = 0 + ep_terminates = 0 + + _EMA = 0.1 + ema_mlp_infer_ms = 0.0 + ema_env_step_ms = 0.0 + + try: + while not stop_event.is_set(): + if actor_weight_sync.version > local_actor_weight_version: + actor_sd = dict(actor.state_dict()) + local_actor_weight_version = actor_weight_sync.read_weights_into(actor_sd) + actor.load_state_dict(actor_sd) + if critic_weight_sync.version > local_critic_weight_version: + critic_sd = dict(critic.state_dict()) + local_critic_weight_version = critic_weight_sync.read_weights_into(critic_sd) + critic.load_state_dict(critic_sd) + + write_buf = storage.write_buffer + for step in range(steps_per_env): + t_mlp = time.perf_counter() + with torch.no_grad(): + obs_torch.copy_(torch.from_numpy(obs_np)) + priv_info_torch.copy_(torch.from_numpy(priv_info_np)) + actions_torch = actor(obs_td, stochastic_output=True) + log_probs_torch = actor.get_output_log_prob(actions_torch) + actions_np = actions_torch.cpu().numpy() + ema_mlp_infer_ms = (1 - _EMA) * ema_mlp_infer_ms + _EMA * ( + (time.perf_counter() - t_mlp) * 1000 + ) + + write_buf["obs"][:, step, :] = obs_np + if critic_np is not None: + write_buf["critic"][:, step, :] = critic_np + write_buf["actions"][:, step, :] = actions_np + write_buf["log_probs"][:, step] = log_probs_torch.cpu().numpy().ravel() + + t_env = time.perf_counter() + state = env.step(actions_np) + ema_env_step_ms = (1 - _EMA) * ema_env_step_ms + _EMA * ( + (time.perf_counter() - t_env) * 1000 + ) + + reward_raw = np.asarray(state.reward, dtype=np.float32).ravel() + terminated_raw = np.asarray(state.terminated, dtype=np.float32).ravel() + truncated_raw = np.asarray(state.truncated, dtype=np.float32).ravel() + combined_done_raw = np.clip(terminated_raw + truncated_raw, 0, 1) + + next_actor_obs_np, next_critic_np, next_priv_info_np = ( + split_hora_obs_with_priv_info( + state.obs, + state.info, + ) + ) + next_actor_obs_np = to_float32_np(next_actor_obs_np) + if next_critic_np is not None: + next_critic_np = to_float32_np(next_critic_np) + if next_priv_info_np is not None: + next_priv_info_np = to_float32_np(next_priv_info_np) + if next_priv_info_np is None: + raise ValueError( + "HORA APPO collector did not receive privileged info after step." + ) + + terminal_contract = resolve_terminal_observation_contract( + next_obs_batch_size=next_actor_obs_np.shape[0], + final_observation=getattr(state, "final_observation", None), + done=combined_done_raw > 0.5, + info=state.info, + truncated=truncated_raw, + ) + terminal_priv_info = None + if ( + terminal_contract.terminal_obs is not None + and terminal_contract.terminal_critic is not None + ): + _, _, terminal_priv_info = split_hora_obs_with_priv_info( + { + "obs": terminal_contract.terminal_obs, + "critic": terminal_contract.terminal_critic, + } + ) + if terminal_priv_info is not None: + terminal_priv_info = to_float32_np(terminal_priv_info) + + reward_raw += compute_hora_timeout_bootstrap_correction( + critic=critic, + collector_device=collector_device, + gamma=float(cfg["algorithm"].get("gamma", 0.99)), + timeout_mask=terminal_contract.timeout_terminal_mask, + final_obs=( + terminal_contract.terminal_obs + if terminal_contract.terminal_obs is not None + else next_actor_obs_np + ), + final_critic=( + terminal_contract.terminal_critic + if terminal_contract.terminal_critic is not None + else next_critic_np + ), + final_priv_info=( + terminal_priv_info if terminal_priv_info is not None else next_priv_info_np + ), + ) + + write_buf["rewards"][:, step] = reward_raw + write_buf["dones"][:, step] = combined_done_raw + write_buf["truncated"][:, step] = truncated_raw + + total_steps += num_envs + current_ep_rewards += reward_raw + current_ep_lengths += 1 + reset_indices = np.where(combined_done_raw > 0.5)[0] + if len(reset_indices) > 0: + ep_rewards.extend(current_ep_rewards[reset_indices].tolist()) + ep_lengths.extend(current_ep_lengths[reset_indices].tolist()) + current_ep_rewards[reset_indices] = 0.0 + current_ep_lengths[reset_indices] = 0 + ep_timeouts += int(np.sum(truncated_raw[reset_indices] > 0.5)) + ep_terminates += int(np.sum(truncated_raw[reset_indices] <= 0.5)) + + log_info = state.info.get("log", {}) + for k, v in log_info.items(): + if k.startswith("reward/"): + ep_reward_components[k].append(v) + + if metrics_queue is not None and total_steps % (num_envs * 10) == 0 and ep_rewards: + try: + msg = { + "total_steps": total_steps, + "mean_ep_reward": statistics.mean(ep_rewards[-100:]), + "mean_ep_length": statistics.mean(ep_lengths[-100:]) + if ep_lengths + else 0.0, + } + total_ep = ep_timeouts + ep_terminates + if total_ep > 0: + msg["timeout_rate"] = ep_timeouts / total_ep + msg["terminated_rate"] = ep_terminates / total_ep + ep_timeouts = 0 + ep_terminates = 0 + msg["collector_timing_ms"] = { + "mlp_infer_ms": ema_mlp_infer_ms, + "env_step_total_ms": ema_env_step_ms, + } + if ep_reward_components: + msg["reward_components"] = { + k: statistics.mean(v) for k, v in ep_reward_components.items() if v + } + ep_reward_components.clear() + metrics_queue.put_nowait(msg) + except Exception as e: + print(f"[HoraAPPOWorker] metrics enqueue error: {e}", file=sys.stderr) + + obs_np = next_actor_obs_np + critic_np = next_critic_np + priv_info_np = next_priv_info_np + + write_buf["last_obs"][:] = obs_np + if critic_np is not None: + write_buf["last_critic"][:] = critic_np + storage.signal_write_done() + + except Exception as e: + import traceback + + print(f"\n[HORA APPO WORKER CRASH]: {e}\n", file=sys.stderr) + traceback.print_exc(file=sys.stderr) + if metrics_queue is not None: + try: + metrics_queue.put_nowait({"error": str(e)}) + except Exception: + pass + stop_event.set() + raise + + storage.close() + actor_weight_sync.close() + critic_weight_sync.close() + env.close() diff --git a/src/unilab/algos/torch/hora/distill.py b/src/unilab/algos/torch/hora/distill.py new file mode 100644 index 000000000..fa66dd5e9 --- /dev/null +++ b/src/unilab/algos/torch/hora/distill.py @@ -0,0 +1,372 @@ +from __future__ import annotations + +import math +import statistics +import time +from collections import deque +from dataclasses import dataclass +from pathlib import Path +from typing import Any, cast + +import torch +from omegaconf import DictConfig, OmegaConf +from tensordict import TensorDict + +from unilab.algos.torch.common.normalization import EmpiricalNormalization +from unilab.algos.torch.hora.models import HoraActorModel, HoraSharedActorCritic + + +@dataclass +class HoraDistillStats: + agent_steps: int = 0 + best_reward: float = float("-inf") + mean_reward: float = float("nan") + mean_episode_length: float = float("nan") + + +def build_student_actor_and_normalizer( + env, + cfg: DictConfig, + *, + device: torch.device, +) -> tuple[HoraActorModel, EmpiricalNormalization]: + actor_obs = env.get_observations() + actor_dim = int(actor_obs["actor"].shape[-1]) + priv_info_dim = int(actor_obs["priv_info"].shape[-1]) + proprio_hist_shape = actor_obs["proprio_hist"].shape[1:] + + model_cfg = OmegaConf.to_container(cfg.algo.model, resolve=True) + assert isinstance(model_cfg, dict) + shared = HoraSharedActorCritic( + obs_dim=actor_dim, + action_dim=int(env.num_actions), + priv_info_dim=priv_info_dim, + actor_hidden_dims=model_cfg.get("hidden_dims", [512, 256, 128]), + activation=model_cfg.get("activation", "elu"), + obs_normalization=model_cfg.get("obs_normalization", True), + distribution_cfg=model_cfg.get("distribution_cfg", {"init_std": 1.0, "std_type": "scalar"}), + priv_info_embed_dim=model_cfg.get("priv_info_embed_dim", priv_info_dim), + priv_mlp_hidden_dims=model_cfg.get("priv_mlp_hidden_dims", [256, 128, 8]), + use_student_encoder=True, + proprio_hist_len=int(proprio_hist_shape[0]), + proprio_frame_dim=int(proprio_hist_shape[1]), + ).to(device) + actor = HoraActorModel( + actor_obs, + {"actor": ["actor"], "critic": ["actor"]}, + "actor", + int(env.num_actions), + shared_model=shared, + use_student_encoder=True, + ).to(device) + hist_normalizer = EmpiricalNormalization(proprio_hist_shape, device=device) + return actor, hist_normalizer + + +def load_teacher_actor_weights( + actor: HoraActorModel, + teacher_checkpoint: str | Path, + *, + teacher_algo_family: str, + device: torch.device, +) -> None: + checkpoint = torch.load(teacher_checkpoint, map_location=device, weights_only=False) + actor_state_key = { + "ppo": "actor_state_dict", + "appo": "actor", + }.get(str(teacher_algo_family)) + if actor_state_key is None: + raise ValueError( + "Unsupported HORA teacher algorithm family for distillation: " + f"{teacher_algo_family!r}. Expected one of ['ppo', 'appo']." + ) + actor_state = checkpoint.get(actor_state_key) + if actor_state is None: + raise ValueError( + "Checkpoint does not contain the expected teacher actor weights. " + f"algo_family={teacher_algo_family!r} expected_key={actor_state_key!r} " + f"checkpoint={teacher_checkpoint}" + ) + actor.load_state_dict(actor_state, strict=False) + + +def load_distilled_checkpoint( + actor: HoraActorModel, + hist_normalizer: EmpiricalNormalization, + checkpoint_path: str | Path, + *, + device: torch.device, +) -> dict[str, Any]: + checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False) + model_state = checkpoint.get("model_state_dict") + if model_state is None: + raise ValueError(f"Checkpoint does not contain model_state_dict: {checkpoint_path}") + actor.load_state_dict(model_state, strict=True) + + history_normalizer = checkpoint.get("history_normalizer") + if history_normalizer is not None: + hist_normalizer.load_state_dict(history_normalizer) + return cast(dict[str, Any], checkpoint) + + +class HoraDistillationTrainer: + """Stage-2 HORA latent distillation trainer.""" + + def __init__( + self, + env, + cfg: DictConfig, + *, + device: str, + log_dir: str | Path, + teacher_checkpoint: str | Path, + teacher_algo_family: str, + teacher_metadata: dict[str, Any] | None = None, + distill_runtime_cfg: DictConfig, + logger, + ) -> None: + self.env = env + self.cfg = cfg + self.device = torch.device(device) + self.log_dir = Path(log_dir) + self.logger = logger + self.teacher_checkpoint = Path(teacher_checkpoint) + self.teacher_algo_family = str(teacher_algo_family) + self.teacher_metadata = dict(teacher_metadata or {}) + self.distill_runtime_cfg = OmegaConf.to_container(distill_runtime_cfg, resolve=True) + self.actor, self.hist_normalizer = build_student_actor_and_normalizer( + env, + cfg, + device=self.device, + ) + self.optimizer = torch.optim.Adam( + self._trainable_parameters(), lr=float(cfg.algo.learning_rate) + ) + self.stats = HoraDistillStats() + self._reward_buffer: deque[float] = deque(maxlen=100) + self._episode_length_buffer: deque[float] = deque(maxlen=100) + self._step_reward = torch.zeros((env.num_envs,), dtype=torch.float32, device=self.device) + self._step_length = torch.zeros((env.num_envs,), dtype=torch.float32, device=self.device) + self._tb_writer = self._build_tensorboard_writer() + self._load_teacher_checkpoint() + + def _trainable_parameters(self) -> list[torch.nn.Parameter]: + params: list[torch.nn.Parameter] = [] + for name, param in self.actor.named_parameters(): + requires_grad = "adapt_tconv" in name + param.requires_grad = requires_grad + if requires_grad: + params.append(param) + return params + + def _load_teacher_checkpoint(self) -> None: + load_teacher_actor_weights( + self.actor, + self.teacher_checkpoint, + teacher_algo_family=self.teacher_algo_family, + device=self.device, + ) + self.actor.train() + self.actor.shared.obs_normalizer.eval() + + def _build_tensorboard_writer(self) -> Any | None: + """Create the stage-2 TensorBoard writer when the config requests it. + + Args: + None. + + Returns: + Summary writer rooted at ``/tb``, or ``None`` when scalar + backend logging is disabled or TensorBoard is unavailable. + """ + logger_type = str(OmegaConf.select(self.cfg, "training.logger", default="tensorboard")) + if logger_type.lower() != "tensorboard": + return None + + try: + from torch.utils.tensorboard import SummaryWriter + except ImportError: + self.logger.warning( + "tensorboard is not installed; disabling HORA distillation TensorBoard logging." + ) + return None + + tb_dir = self.log_dir / "tb" + tb_dir.mkdir(parents=True, exist_ok=True) + self.logger.info("TensorBoard: %s", tb_dir) + return SummaryWriter(log_dir=str(tb_dir)) + + def _add_scalar_if_finite(self, tag: str, value: float, *, step: int) -> None: + """Write a scalar only when a TensorBoard writer exists and the value is finite. + + Args: + tag: TensorBoard metric name. + value: Scalar value to record. + step: Global step associated with the scalar. + + Returns: + None. Invalid or disabled values are skipped silently so early NaNs + from unfinished episodes do not pollute the event stream. + """ + if self._tb_writer is None or not math.isfinite(value): + return + self._tb_writer.add_scalar(tag, value, step) + + def _log_tensorboard_step(self, *, loss: float, elapsed: float) -> None: + """Record the latest distillation scalars to TensorBoard. + + Args: + loss: Latest latent-distillation loss. + elapsed: Wall-clock training time since the run started. + + Returns: + None. Metrics are written at the current agent-step count. + """ + if self._tb_writer is None: + return + + step = self.stats.agent_steps + self._add_scalar_if_finite("train/loss", loss, step=step) + self._add_scalar_if_finite("reward/mean", self.stats.mean_reward, step=step) + self._add_scalar_if_finite("reward/best", self.stats.best_reward, step=step) + self._add_scalar_if_finite( + "episode/length", + self.stats.mean_episode_length, + step=step, + ) + self._add_scalar_if_finite("perf/fps", step / max(elapsed, 1e-6), step=step) + self._add_scalar_if_finite("perf/training_time_sec", elapsed, step=step) + self._tb_writer.flush() + + def _normalize_student_obs(self, obs_td) -> dict[str, torch.Tensor]: + actor_obs = obs_td["actor"].to(self.device) + proprio_hist = obs_td["proprio_hist"].to(self.device) + return { + "actor": actor_obs, + "priv_info": obs_td["priv_info"].to(self.device), + "proprio_hist": self.hist_normalizer(proprio_hist), + } + + @staticmethod + def _next_interval_boundary(current_steps: int, interval_steps: int) -> int | None: + """Return the next positive save boundary after the current step count. + + Args: + current_steps: Number of agent steps already completed. + interval_steps: Positive interval in agent steps between saves. + + Returns: + The next interval boundary, or ``None`` when periodic saving is disabled. + """ + if interval_steps <= 0: + return None + return ((current_steps // interval_steps) + 1) * interval_steps + + def train(self) -> None: + obs_td, _ = self.env.reset() + max_agent_steps = int(self.cfg.algo.max_agent_steps) + save_interval = int(self.cfg.algo.save_interval_steps) + log_interval = int(self.cfg.algo.log_interval_steps) + next_log_steps = self._next_interval_boundary(self.stats.agent_steps, log_interval) + next_save_steps = self._next_interval_boundary(self.stats.agent_steps, save_interval) + start_time = time.time() + last_loss = float("nan") + + try: + while self.stats.agent_steps < max_agent_steps: + norm_obs = self._normalize_student_obs(obs_td) + obs_batch = { + key: value.detach() if key == "actor" else value + for key, value in norm_obs.items() + } + td = TensorDict(obs_batch, batch_size=obs_td.batch_size, device=self.device) + _, core_output = self.actor.shared.policy_mean(td, prefer_student=True) + loss = torch.mean( + (core_output.privileged_latent - core_output.privileged_target.detach()) ** 2 + ) + last_loss = float(loss.item()) + + self.optimizer.zero_grad() + loss.backward() + self.optimizer.step() + + with torch.no_grad(): + actions = self.actor(td, stochastic_output=False).clamp_(-1.0, 1.0) + obs_td, rewards, dones, infos = self.env.step(actions) + rewards = rewards.to(self.device) + dones = dones.to(self.device) + self.stats.agent_steps += int(self.env.num_envs) + + self._step_reward += rewards + self._step_length += 1 + done_idx = torch.nonzero(dones, as_tuple=False).flatten() + if len(done_idx) > 0: + completed_rewards = self._step_reward[done_idx] + completed_lengths = self._step_length[done_idx] + done_mean_reward = float(torch.mean(completed_rewards).item()) + self._reward_buffer.extend(completed_rewards.detach().cpu().numpy().tolist()) + self._episode_length_buffer.extend( + completed_lengths.detach().cpu().numpy().tolist() + ) + self.stats.mean_reward = float(statistics.mean(self._reward_buffer)) + self.stats.mean_episode_length = float( + statistics.mean(self._episode_length_buffer) + ) + self.stats.best_reward = max(self.stats.best_reward, done_mean_reward) + self._step_reward[done_idx] = 0.0 + self._step_length[done_idx] = 0.0 + + if next_log_steps is not None and self.stats.agent_steps >= next_log_steps: + elapsed = max(time.time() - start_time, 1e-6) + self.logger.info( + "agent_steps=%d loss=%.6f mean_reward=%.4f best_reward=%.4f " + "mean_episode_length=%.2f training_time=%.2fs fps=%.1f", + self.stats.agent_steps, + last_loss, + self.stats.mean_reward, + self.stats.best_reward, + self.stats.mean_episode_length, + elapsed, + self.stats.agent_steps / elapsed, + ) + self._log_tensorboard_step(loss=last_loss, elapsed=elapsed) + next_log_steps = self._next_interval_boundary( + self.stats.agent_steps, log_interval + ) + + if next_save_steps is not None and self.stats.agent_steps >= next_save_steps: + self.save(self.log_dir / f"hora_stage2_{self.stats.agent_steps}.pt") + next_save_steps = self._next_interval_boundary( + self.stats.agent_steps, save_interval + ) + + self.save(self.log_dir / "hora_stage2_last.pt") + total_elapsed = max(time.time() - start_time, 1e-6) + self.logger.info( + "training_complete agent_steps=%d mean_reward=%.4f best_reward=%.4f " + "mean_episode_length=%.2f training_time=%.2fs", + self.stats.agent_steps, + self.stats.mean_reward, + self.stats.best_reward, + self.stats.mean_episode_length, + total_elapsed, + ) + self._log_tensorboard_step(loss=last_loss, elapsed=total_elapsed) + finally: + if self._tb_writer is not None: + self._tb_writer.close() + self._tb_writer = None + + def save(self, path: str | Path) -> None: + torch.save( + { + "model_state_dict": self.actor.state_dict(), + "history_normalizer": self.hist_normalizer.state_dict(), + "agent_steps": self.stats.agent_steps, + "teacher_checkpoint": str(self.teacher_checkpoint), + "teacher_algo_family": self.teacher_algo_family, + "teacher_metadata": self.teacher_metadata, + "distill_runtime_cfg": self.distill_runtime_cfg, + }, + Path(path), + ) diff --git a/src/unilab/algos/torch/hora/distill_config.py b/src/unilab/algos/torch/hora/distill_config.py new file mode 100644 index 000000000..71f3bf595 --- /dev/null +++ b/src/unilab/algos/torch/hora/distill_config.py @@ -0,0 +1,214 @@ +"""HORA distillation config and teacher-owner resolution helpers.""" + +from __future__ import annotations + +import re +from pathlib import Path +from typing import Any, cast + +from omegaconf import DictConfig, OmegaConf + +from unilab.training.run import resolve_task_checkpoint_path + +_REPO_ROOT = Path(__file__).resolve().parents[5] + + +def _root(root_dir: str | Path | None) -> Path: + return Path(root_dir) if root_dir is not None else _REPO_ROOT + + +def _load_yaml_config(path: Path) -> DictConfig: + loaded = OmegaConf.load(path) + if not isinstance(loaded, DictConfig): + raise TypeError(f"Expected DictConfig from {path}, got {type(loaded)!r}") + return loaded + + +def _sanitize_path_token(value: str, *, fallback: str) -> str: + sanitized = re.sub(r"[^A-Za-z0-9._-]+", "-", str(value)).strip("-._") + return sanitized or fallback + + +def load_teacher_owner_config( + algo_family: str, + task: str, + *, + root_dir: str | Path | None = None, +) -> DictConfig: + """Load a HORA teacher owner config and its direct owner defaults.""" + root = _root(root_dir) + owner_path = root / "conf" / str(algo_family) / "task" / f"{task}.yaml" + owner_cfg = _load_yaml_config(owner_path) + merged_cfg = OmegaConf.create() + for default_entry in owner_cfg.get("defaults", []): + if not isinstance(default_entry, str) or default_entry == "_self_": + continue + include_path = root / "conf" / str(algo_family) / f"{default_entry.lstrip('/')}.yaml" + merged_cfg = OmegaConf.merge(merged_cfg, _load_yaml_config(include_path)) + return cast(DictConfig, OmegaConf.merge(merged_cfg, owner_cfg)) + + +def get_teacher_owner_spec(cfg: DictConfig) -> tuple[str | None, str | None]: + """Resolve the teacher algo family and task owner from distillation config.""" + algo_family = OmegaConf.select(cfg, "teacher.algo_family") + task = OmegaConf.select(cfg, "teacher.task") + if algo_family in (None, "") or task in (None, ""): + return None, None + return str(algo_family), str(task) + + +def teacher_default_cfg( + cfg: DictConfig, + *, + root_dir: str | Path | None = None, +) -> DictConfig: + """Build HORA student defaults from the selected teacher owner YAML.""" + teacher_algo_family, teacher_task = get_teacher_owner_spec(cfg) + if teacher_algo_family is None or teacher_task is None: + return OmegaConf.create() + + teacher_cfg = load_teacher_owner_config( + teacher_algo_family, + teacher_task, + root_dir=root_dir, + ) + actor_cfg = OmegaConf.to_container(OmegaConf.select(teacher_cfg, "algo.actor"), resolve=True) + if not isinstance(actor_cfg, dict): + actor_cfg = {} + actor_cfg = dict(actor_cfg) + actor_class_name = str(actor_cfg.get("class_name", "")) + if "HoraActorModel" not in actor_class_name: + raise ValueError( + "HORA distillation teacher owner must resolve to HoraActorModel. " + f"Got algo_family={teacher_algo_family} task={teacher_task} " + f"actor.class_name={actor_class_name!r}." + ) + actor_cfg.pop("class_name", None) + distribution_cfg = actor_cfg.get("distribution_cfg") + if isinstance(distribution_cfg, dict): + distribution_cfg = { + key: value for key, value in distribution_cfg.items() if key != "class_name" + } + + return OmegaConf.create( + { + "training": OmegaConf.select(teacher_cfg, "training"), + "reward": OmegaConf.select(teacher_cfg, "reward"), + "env": OmegaConf.select(teacher_cfg, "env"), + "algo": { + "model": { + "hidden_dims": actor_cfg.get("hidden_dims"), + "activation": actor_cfg.get("activation"), + "obs_normalization": actor_cfg.get("obs_normalization"), + "priv_info_embed_dim": actor_cfg.get("priv_info_embed_dim"), + "priv_mlp_hidden_dims": actor_cfg.get("priv_mlp_hidden_dims"), + "distribution_cfg": distribution_cfg, + } + }, + } + ) + + +def apply_teacher_defaults( + cfg: DictConfig, + *, + root_dir: str | Path | None = None, +) -> DictConfig: + """Merge teacher-owner defaults under the user distillation config.""" + return cast(DictConfig, OmegaConf.merge(teacher_default_cfg(cfg, root_dir=root_dir), cfg)) + + +def resolved_distill_runtime_cfg(cfg: DictConfig) -> DictConfig: + """Return stage-2 playback fields that do not depend on teacher algorithm.""" + model_cfg = OmegaConf.select(cfg, "algo.model") + return OmegaConf.create( + { + "training": { + "task_name": OmegaConf.select(cfg, "training.task_name"), + "sim_backend": OmegaConf.select(cfg, "training.sim_backend"), + "render_spacing": OmegaConf.select(cfg, "training.render_spacing"), + "cam_distance": OmegaConf.select(cfg, "training.cam_distance"), + "cam_elevation": OmegaConf.select(cfg, "training.cam_elevation"), + "cam_azimuth": OmegaConf.select(cfg, "training.cam_azimuth"), + "cam_lookat": OmegaConf.select(cfg, "training.cam_lookat"), + "cam_tracking": OmegaConf.select(cfg, "training.cam_tracking"), + "cam_tracking_env_idx": OmegaConf.select(cfg, "training.cam_tracking_env_idx"), + "cam_tracking_extra_envs": OmegaConf.select( + cfg, "training.cam_tracking_extra_envs" + ), + }, + "reward": OmegaConf.select(cfg, "reward"), + "env": OmegaConf.select(cfg, "env"), + "algo": { + "model": ( + OmegaConf.to_container(model_cfg, resolve=True) if model_cfg is not None else {} + ) + }, + } + ) + + +def teacher_run_metadata( + cfg: DictConfig, + *, + teacher_algo_family: str, + teacher_checkpoint: Path, + root_dir: str | Path | None = None, +) -> dict[str, Any]: + """Build explicit teacher provenance metadata for distillation outputs.""" + teacher_task = OmegaConf.select(cfg, "teacher.task") + root = _root(root_dir).resolve() + checkpoint_path = teacher_checkpoint.resolve() + try: + checkpoint_display = str(checkpoint_path.relative_to(root)) + except ValueError: + checkpoint_display = str(checkpoint_path) + + checkpoint_name = checkpoint_path.name + return { + "algo_family": str(teacher_algo_family), + "task": None if teacher_task in (None, "") else str(teacher_task), + "checkpoint_path": checkpoint_display, + "checkpoint_name": checkpoint_name, + "checkpoint_stem": checkpoint_path.stem, + "run_name": checkpoint_path.parent.name, + "run_slug": f"teacher-{_sanitize_path_token(teacher_algo_family, fallback='teacher')}", + } + + +def resolve_teacher_checkpoint_path( + cfg: DictConfig, + *, + root_dir: str | Path | None = None, +) -> tuple[Path | None, Path | None]: + """Resolve the selected HORA teacher checkpoint through owner metadata.""" + teacher_algo_family, teacher_task = get_teacher_owner_spec(cfg) + if teacher_algo_family is None or teacher_task is None: + return None, None + + root = _root(root_dir) + teacher_cfg = load_teacher_owner_config( + teacher_algo_family, + teacher_task, + root_dir=root, + ) + teacher_task_name = OmegaConf.select(teacher_cfg, "training.task_name") + teacher_algo_log_name = OmegaConf.select(teacher_cfg, "algo.algo_log_name") + if teacher_task_name in (None, "") or teacher_algo_log_name in (None, ""): + raise ValueError( + "Teacher owner config must define training.task_name and algo.algo_log_name. " + f"Got algo_family={teacher_algo_family} task={teacher_task}." + ) + + selected_checkpoint = OmegaConf.select(cfg, "algo.checkpoint", default=-1) + return resolve_task_checkpoint_path( + root, + task_name=str(teacher_task_name), + load_run=str(OmegaConf.select(cfg, "algo.load_run", default="-1")), + algo_log_name=str(teacher_algo_log_name), + checkpoint=( + str(selected_checkpoint) if selected_checkpoint not in (None, "", -1, "-1") else None + ), + suffix=".pt", + log_root=OmegaConf.select(cfg, "training.log_root"), + ) diff --git a/src/unilab/algos/torch/hora/models.py b/src/unilab/algos/torch/hora/models.py new file mode 100644 index 000000000..789e88780 --- /dev/null +++ b/src/unilab/algos/torch/hora/models.py @@ -0,0 +1,454 @@ +from __future__ import annotations + +import copy +from dataclasses import dataclass +from typing import Any, cast + +import torch +import torch.nn as nn +from rsl_rl.modules import EmpiricalNormalization, GaussianDistribution +from tensordict import TensorDict + + +def _build_activation(name: str) -> nn.Module: + normalized = str(name).strip().lower() + if normalized == "elu": + return nn.ELU() + if normalized == "relu": + return nn.ReLU() + if normalized == "tanh": + return nn.Tanh() + raise ValueError(f"Unsupported activation: {name!r}") + + +class _MLP(nn.Module): + def __init__( + self, input_dim: int, hidden_dims: list[int] | tuple[int, ...], activation: str + ) -> None: + super().__init__() + layers: list[nn.Module] = [] + current_dim = input_dim + for hidden_dim in hidden_dims: + layers.append(nn.Linear(current_dim, int(hidden_dim))) + layers.append(_build_activation(activation)) + current_dim = int(hidden_dim) + self.net = nn.Sequential(*layers) + self.output_dim = current_dim + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.net(x) + + +class ProprioAdaptTConv(nn.Module): + """Temporal adaptation encoder used by HORA stage-2 distillation.""" + + def __init__(self, frame_dim: int, latent_dim: int) -> None: + super().__init__() + self.channel_transform = nn.Sequential( + nn.Linear(frame_dim, frame_dim), + nn.ReLU(inplace=True), + nn.Linear(frame_dim, frame_dim), + nn.ReLU(inplace=True), + ) + self.temporal_aggregation = nn.Sequential( + nn.Conv1d(frame_dim, frame_dim, kernel_size=9, stride=2), + nn.ReLU(inplace=True), + nn.Conv1d(frame_dim, frame_dim, kernel_size=5, stride=1), + nn.ReLU(inplace=True), + nn.Conv1d(frame_dim, frame_dim, kernel_size=5, stride=1), + nn.ReLU(inplace=True), + ) + self.low_dim_proj = nn.Linear(frame_dim * 3, latent_dim) + self._init_weights() + + def _init_weights(self) -> None: + for module in self.modules(): + if isinstance(module, nn.Conv1d): + fan_out = module.kernel_size[0] * module.out_channels + module.weight.data.normal_(mean=0.0, std=(2.0 / fan_out) ** 0.5) + if module.bias is not None: + nn.init.zeros_(module.bias) + if isinstance(module, nn.Linear) and module.bias is not None: + nn.init.zeros_(module.bias) + + def forward(self, proprio_hist: torch.Tensor) -> torch.Tensor: + x = self.channel_transform(proprio_hist) + x = x.permute(0, 2, 1) + x = self.temporal_aggregation(x) + return self.low_dim_proj(x.flatten(1)) + + +@dataclass +class HoraCoreOutput: + policy_obs: torch.Tensor + trunk_latent: torch.Tensor + privileged_latent: torch.Tensor + privileged_target: torch.Tensor + + +class HoraSharedActorCritic(nn.Module): + """Shared-backbone HORA actor-critic with optional adaptation encoder.""" + + def __init__( + self, + obs_dim: int, + action_dim: int, + *, + priv_info_dim: int, + priv_info_embed_dim: int = 8, + actor_hidden_dims: list[int] | tuple[int, ...] = (512, 256, 128), + priv_mlp_hidden_dims: list[int] | tuple[int, ...] = (256, 128, 8), + activation: str = "elu", + obs_normalization: bool = False, + distribution_cfg: dict[str, Any] | None = None, + use_student_encoder: bool = False, + proprio_hist_len: int = 30, + proprio_frame_dim: int | None = None, + ) -> None: + super().__init__() + self.obs_dim = int(obs_dim) + self.action_dim = int(action_dim) + self.priv_info_dim = int(priv_info_dim) + self.priv_info_embed_dim = int(priv_info_embed_dim) + self.use_student_encoder = bool(use_student_encoder) + self.proprio_hist_len = int(proprio_hist_len) + self.proprio_frame_dim = ( + int(proprio_frame_dim) if proprio_frame_dim is not None else self.obs_dim // 3 + ) + + self.obs_normalizer = ( + EmpiricalNormalization(self.obs_dim) if obs_normalization else nn.Identity() + ) + self.priv_encoder = _MLP(self.priv_info_dim, list(priv_mlp_hidden_dims), activation) + self.trunk = _MLP( + self.obs_dim + self.priv_info_embed_dim, list(actor_hidden_dims), activation + ) + self.value_head = nn.Linear(self.trunk.output_dim, 1) + self.mu_head = nn.Linear(self.trunk.output_dim, self.action_dim) + self.distribution = GaussianDistribution( + self.action_dim, + **( + { + key: value + for key, value in ( + distribution_cfg + if distribution_cfg is not None + else {"init_std": 1.0, "std_type": "scalar"} + ).items() + if key != "class_name" + } + ), + ) + self.adapt_tconv = ( + ProprioAdaptTConv(self.proprio_frame_dim, self.priv_info_embed_dim) + if self.use_student_encoder + else None + ) + self._init_linear_biases() + + def _init_linear_biases(self) -> None: + for module in self.modules(): + if isinstance(module, nn.Linear) and module.bias is not None: + nn.init.zeros_(module.bias) + + def _normalize_actor_obs(self, actor_obs: torch.Tensor) -> torch.Tensor: + return self.obs_normalizer(actor_obs) + + def update_normalization(self, obs: TensorDict) -> None: + if isinstance(self.obs_normalizer, EmpiricalNormalization): + self.obs_normalizer.update(obs["actor"]) + + def _zero_privileged_latent( + self, batch_size: int, device: torch.device, dtype: torch.dtype + ) -> torch.Tensor: + return torch.zeros((batch_size, self.priv_info_embed_dim), device=device, dtype=dtype) + + def encode_privileged_info(self, priv_info: torch.Tensor | None) -> torch.Tensor: + if priv_info is None: + raise ValueError("priv_info is required to compute the HORA teacher latent") + return torch.tanh(self.priv_encoder(priv_info)) + + def encode_proprio_history(self, proprio_hist: torch.Tensor) -> torch.Tensor: + if self.adapt_tconv is None: + raise RuntimeError("HORA adaptation encoder is not enabled") + return torch.tanh(self.adapt_tconv(proprio_hist)) + + def build_core_output(self, obs: TensorDict, *, prefer_student: bool) -> HoraCoreOutput: + actor_obs = obs["actor"] + policy_obs = self._normalize_actor_obs(actor_obs) + priv_info = obs.get("priv_info") + proprio_hist = obs.get("proprio_hist") + + privileged_target = ( + self.encode_privileged_info(priv_info) + if priv_info is not None + else self._zero_privileged_latent(actor_obs.shape[0], actor_obs.device, actor_obs.dtype) + ) + + if prefer_student and self.adapt_tconv is not None and proprio_hist is not None: + privileged_latent = self.encode_proprio_history(proprio_hist) + elif priv_info is not None: + privileged_latent = privileged_target + else: + privileged_latent = self._zero_privileged_latent( + actor_obs.shape[0], actor_obs.device, actor_obs.dtype + ) + + trunk_input = torch.cat([policy_obs, privileged_latent], dim=-1) + trunk_latent = self.trunk(trunk_input) + return HoraCoreOutput( + policy_obs=policy_obs, + trunk_latent=trunk_latent, + privileged_latent=privileged_latent, + privileged_target=privileged_target, + ) + + def policy_mean( + self, obs: TensorDict, *, prefer_student: bool + ) -> tuple[torch.Tensor, HoraCoreOutput]: + core_output = self.build_core_output(obs, prefer_student=prefer_student) + return self.mu_head(core_output.trunk_latent), core_output + + def value( + self, obs: TensorDict, *, prefer_student: bool + ) -> tuple[torch.Tensor, HoraCoreOutput]: + core_output = self.build_core_output(obs, prefer_student=prefer_student) + return self.value_head(core_output.trunk_latent), core_output + + +class _HoraInferenceModule(nn.Module): + input_names = ["actor", "priv_info", "proprio_hist"] + output_names = ["actions"] + + def __init__( + self, + *, + obs_normalizer: nn.Module, + priv_encoder: nn.Module, + trunk: nn.Module, + mu_head: nn.Module, + obs_dim: int, + priv_info_dim: int, + proprio_hist_len: int, + proprio_frame_dim: int, + verbose: bool = False, + adapt_tconv: nn.Module | None = None, + prefer_student: bool = False, + ) -> None: + super().__init__() + self.obs_normalizer = obs_normalizer + self.priv_encoder = priv_encoder + self.trunk = trunk + self.mu_head = mu_head + self.adapt_tconv = adapt_tconv + self.prefer_student = bool(prefer_student) + self.obs_dim = int(obs_dim) + self.priv_info_dim = int(priv_info_dim) + self.proprio_hist_len = int(proprio_hist_len) + self.proprio_frame_dim = int(proprio_frame_dim) + self.verbose = bool(verbose) + + def forward( + self, actor: torch.Tensor, priv_info: torch.Tensor, proprio_hist: torch.Tensor + ) -> torch.Tensor: + policy_obs = self.obs_normalizer(actor) + if self.prefer_student: + if self.adapt_tconv is None: + raise RuntimeError("HORA adaptation encoder export requires adapt_tconv") + privileged_latent = torch.tanh(self.adapt_tconv(proprio_hist)) + else: + privileged_latent = torch.tanh(self.priv_encoder(priv_info)) + trunk_input = torch.cat([policy_obs, privileged_latent], dim=-1) + return self.mu_head(self.trunk(trunk_input)) + + def get_dummy_inputs(self) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + return ( + torch.zeros(1, self.obs_dim), + torch.zeros(1, self.priv_info_dim), + torch.zeros(1, self.proprio_hist_len, self.proprio_frame_dim), + ) + + +class HoraActorModel(nn.Module): + is_recurrent: bool = False + + def __init__( + self, + obs: TensorDict, + obs_groups: dict[str, list[str]], + obs_set: str, + output_dim: int, + *, + shared_model: HoraSharedActorCritic | None = None, + hidden_dims: list[int] | tuple[int, ...] = (512, 256, 128), + activation: str = "elu", + obs_normalization: bool = False, + distribution_cfg: dict[str, Any] | None = None, + priv_info_dim: int | None = None, + priv_info_embed_dim: int = 8, + priv_mlp_hidden_dims: list[int] | tuple[int, ...] = (256, 128, 8), + use_student_encoder: bool = False, + proprio_hist_len: int = 30, + proprio_frame_dim: int | None = None, + ) -> None: + del obs_groups, obs_set + super().__init__() + if shared_model is None: + shared_model = HoraSharedActorCritic( + obs_dim=int(obs["actor"].shape[-1]), + action_dim=output_dim, + priv_info_dim=int( + priv_info_dim if priv_info_dim is not None else obs.get("priv_info").shape[-1] + ), + priv_info_embed_dim=priv_info_embed_dim, + actor_hidden_dims=hidden_dims, + priv_mlp_hidden_dims=priv_mlp_hidden_dims, + activation=activation, + obs_normalization=obs_normalization, + distribution_cfg=distribution_cfg, + use_student_encoder=use_student_encoder, + proprio_hist_len=proprio_hist_len, + proprio_frame_dim=proprio_frame_dim + if proprio_frame_dim is not None + else (int(obs["proprio_hist"].shape[-1]) if "proprio_hist" in obs else None), + ) + self.shared = shared_model + self.prefer_student = bool(use_student_encoder) + + def forward( + self, + obs: TensorDict, + masks: torch.Tensor | None = None, + hidden_state=None, + stochastic_output: bool = False, + ) -> torch.Tensor: + del masks, hidden_state + mean, _ = self.shared.policy_mean(obs, prefer_student=self.prefer_student) + self.shared.distribution.update(mean) + if stochastic_output: + return self.shared.distribution.sample() + return self.shared.distribution.deterministic_output(mean) + + def reset(self, dones: torch.Tensor | None = None, hidden_state=None) -> None: + del dones, hidden_state + + def get_hidden_state(self): + return None + + def detach_hidden_state(self, dones: torch.Tensor | None = None) -> None: + del dones + + @property + def output_mean(self) -> torch.Tensor: + return self.shared.distribution.mean + + @property + def output_std(self) -> torch.Tensor: + return self.shared.distribution.std + + @property + def output_entropy(self) -> torch.Tensor: + return self.shared.distribution.entropy + + @property + def output_distribution_params(self) -> tuple[torch.Tensor, ...]: + return cast(tuple[torch.Tensor, ...], self.shared.distribution.params) + + def get_output_log_prob(self, outputs: torch.Tensor) -> torch.Tensor: + return self.shared.distribution.log_prob(outputs) + + def get_kl_divergence( + self, old_params: tuple[torch.Tensor, ...], new_params: tuple[torch.Tensor, ...] + ) -> torch.Tensor: + return self.shared.distribution.kl_divergence(old_params, new_params) + + def update_normalization(self, obs: TensorDict) -> None: + self.shared.update_normalization(obs) + + def as_jit(self) -> nn.Module: + return _HoraInferenceModule( + obs_normalizer=copy.deepcopy(self.shared.obs_normalizer), + priv_encoder=copy.deepcopy(self.shared.priv_encoder), + trunk=copy.deepcopy(self.shared.trunk), + mu_head=copy.deepcopy(self.shared.mu_head), + obs_dim=self.shared.obs_dim, + priv_info_dim=self.shared.priv_info_dim, + proprio_hist_len=self.shared.proprio_hist_len, + proprio_frame_dim=self.shared.proprio_frame_dim, + adapt_tconv=copy.deepcopy(self.shared.adapt_tconv), + prefer_student=self.prefer_student, + ) + + def as_onnx(self, verbose: bool) -> nn.Module: + return _HoraInferenceModule( + obs_normalizer=copy.deepcopy(self.shared.obs_normalizer), + priv_encoder=copy.deepcopy(self.shared.priv_encoder), + trunk=copy.deepcopy(self.shared.trunk), + mu_head=copy.deepcopy(self.shared.mu_head), + obs_dim=self.shared.obs_dim, + priv_info_dim=self.shared.priv_info_dim, + proprio_hist_len=self.shared.proprio_hist_len, + proprio_frame_dim=self.shared.proprio_frame_dim, + verbose=verbose, + adapt_tconv=copy.deepcopy(self.shared.adapt_tconv), + prefer_student=self.prefer_student, + ) + + +class HoraCriticModel(nn.Module): + is_recurrent: bool = False + + def __init__( + self, + obs: TensorDict, + obs_groups: dict[str, list[str]], + obs_set: str, + output_dim: int, + *, + shared_model: HoraSharedActorCritic | None = None, + hidden_dims: list[int] | tuple[int, ...] = (512, 256, 128), + activation: str = "elu", + obs_normalization: bool = False, + priv_info_dim: int | None = None, + priv_info_embed_dim: int = 8, + priv_mlp_hidden_dims: list[int] | tuple[int, ...] = (256, 128, 8), + ) -> None: + del obs_groups, obs_set, output_dim + super().__init__() + if shared_model is None: + shared_model = HoraSharedActorCritic( + obs_dim=int(obs["actor"].shape[-1]), + action_dim=1, + priv_info_dim=int( + priv_info_dim if priv_info_dim is not None else obs.get("priv_info").shape[-1] + ), + priv_info_embed_dim=priv_info_embed_dim, + actor_hidden_dims=hidden_dims, + priv_mlp_hidden_dims=priv_mlp_hidden_dims, + activation=activation, + obs_normalization=obs_normalization, + ) + self.shared = shared_model + + def forward( + self, + obs: TensorDict, + masks: torch.Tensor | None = None, + hidden_state=None, + stochastic_output: bool = False, + ) -> torch.Tensor: + del masks, hidden_state, stochastic_output + value, _ = self.shared.value(obs, prefer_student=False) + return value + + def reset(self, dones: torch.Tensor | None = None, hidden_state=None) -> None: + del dones, hidden_state + + def get_hidden_state(self): + return None + + def detach_hidden_state(self, dones: torch.Tensor | None = None) -> None: + del dones + + def update_normalization(self, obs: TensorDict) -> None: + del obs diff --git a/src/unilab/algos/torch/hora/observations.py b/src/unilab/algos/torch/hora/observations.py new file mode 100644 index 000000000..5530d801a --- /dev/null +++ b/src/unilab/algos/torch/hora/observations.py @@ -0,0 +1,127 @@ +"""HORA-owned observation helpers for teacher-policy runtime code.""" + +from __future__ import annotations + +from typing import Any + +import numpy as np +from tensordict import TensorDict + +from unilab.utils.tensor import to_torch + + +def split_hora_obs_with_priv_info( + obs: dict[str, np.ndarray], + info: dict[str, Any] | None = None, +) -> tuple[np.ndarray, np.ndarray | None, np.ndarray | None]: + """Split HORA env outputs into actor obs, critic obs, and privileged info. + + Args: + obs: Environment observation dict following the UniLab env contract. + info: Optional env info dict. When present, ``info["critic_info"]`` is the + preferred source of HORA privileged info. + + Returns: + Tuple ``(actor_obs, critic_obs, priv_info)``. ``priv_info`` falls back to the + extra tail of ``critic_obs`` when no explicit ``critic_info`` is provided. + """ + actor_obs = obs["obs"] + critic_obs = obs.get("critic", actor_obs) + + priv_info: np.ndarray | None = None + if isinstance(info, dict): + candidate = info.get("critic_info") + if isinstance(candidate, np.ndarray) and candidate.shape[0] == actor_obs.shape[0]: + priv_info = candidate + + if ( + priv_info is None + and critic_obs is not None + and critic_obs.ndim == 2 + and actor_obs.ndim == 2 + and critic_obs.shape[0] == actor_obs.shape[0] + and critic_obs.shape[1] > actor_obs.shape[1] + ): + priv_info = critic_obs[:, actor_obs.shape[1] :] + + return actor_obs, critic_obs, priv_info + + +def extract_hora_proprio_hist(info: dict[str, Any] | None) -> np.ndarray | None: + """Return HORA proprio-history payload from env info when available. + + Args: + info: Optional env info dict produced by the HORA environment. + + Returns: + Proprio-history array when present, otherwise ``None``. + """ + if not isinstance(info, dict): + return None + proprio_hist = info.get("proprio_hist") + return proprio_hist if isinstance(proprio_hist, np.ndarray) else None + + +def build_hora_obs_tensordict( + obs: dict[str, np.ndarray], + *, + info: dict[str, Any] | None, + device: str, + batch_size: int, + policy_obs: np.ndarray, +) -> TensorDict: + """Build the HORA PPO/APPO observation TensorDict for teacher-policy runtime. + + Args: + obs: Environment observation dict following the UniLab env contract. + info: Optional env info dict containing HORA privileged payloads. + device: Torch device string used for the returned tensors. + batch_size: Number of vectorized environments represented by this batch. + policy_obs: Policy observation array already resolved by the caller. + + Returns: + TensorDict with generic keys plus HORA-specific ``priv_info`` and optional + ``proprio_hist`` when the environment provided them. + """ + actor_obs_np, critic_obs_np, priv_info_np = split_hora_obs_with_priv_info(obs, info) + td_dict = { + "actor": to_torch(actor_obs_np, device), + "policy": to_torch(policy_obs, device), + } + if critic_obs_np is not None: + td_dict["critic"] = to_torch(critic_obs_np, device) + if priv_info_np is not None: + td_dict["priv_info"] = to_torch(priv_info_np, device) + proprio_hist = extract_hora_proprio_hist(info) + if proprio_hist is not None: + td_dict["proprio_hist"] = to_torch(proprio_hist, device) + return TensorDict(td_dict, batch_size=batch_size, device=device) + + +def build_hora_actor_tensordict( + actor_obs: np.ndarray, + *, + priv_info: np.ndarray, + device: str, + batch_size: int, +) -> TensorDict: + """Build the minimal HORA actor TensorDict for APPO play/inference. + + Args: + actor_obs: Actor observation array with shape ``(batch, obs_dim)``. + priv_info: Privileged-info array with shape ``(batch, priv_dim)``. + device: Torch device string used for the returned tensors. + batch_size: Number of vectorized environments represented by this batch. + + Returns: + TensorDict containing grouped HORA actor inputs required by teacher-policy + inference. + """ + return TensorDict( + { + "actor": to_torch(actor_obs, device), + "priv_info": to_torch(priv_info, device), + }, + batch_size=batch_size, + device=device, + ) diff --git a/src/unilab/algos/torch/hora/ppo.py b/src/unilab/algos/torch/hora/ppo.py new file mode 100644 index 000000000..a21909167 --- /dev/null +++ b/src/unilab/algos/torch/hora/ppo.py @@ -0,0 +1,225 @@ +from __future__ import annotations + +from collections.abc import Callable +from itertools import chain +from typing import Any, cast + +import torch +import torch.optim as optim +from rsl_rl.algorithms.ppo import PPO +from rsl_rl.env import VecEnv +from rsl_rl.extensions import resolve_rnd_config, resolve_symmetry_config +from rsl_rl.storage import RolloutStorage +from rsl_rl.utils import resolve_obs_groups, resolve_optimizer +from tensordict import TensorDict + +from unilab.algos.torch.hora.models import HoraActorModel, HoraCriticModel, HoraSharedActorCritic +from unilab.algos.torch.rsl_rl_ppo import FinalObservationAwarePPO + + +class HoraPPO(FinalObservationAwarePPO): + """PPO variant that constructs a shared HORA actor-critic backbone.""" + + def __init__( + self, + actor: HoraActorModel, + critic: HoraCriticModel, + storage: RolloutStorage, + num_learning_epochs: int = 5, + num_mini_batches: int = 4, + clip_param: float = 0.2, + gamma: float = 0.99, + lam: float = 0.95, + value_loss_coef: float = 1.0, + entropy_coef: float = 0.01, + learning_rate: float = 0.001, + max_grad_norm: float = 1.0, + optimizer: str = "adam", + use_clipped_value_loss: bool = True, + schedule: str = "adaptive", + desired_kl: float = 0.01, + normalize_advantage_per_mini_batch: bool = False, + device: str = "cpu", + rnd_cfg: dict | None = None, + symmetry_cfg: dict | None = None, + multi_gpu_cfg: dict | None = None, + ) -> None: + self.device = device + self.is_multi_gpu = multi_gpu_cfg is not None + if multi_gpu_cfg is not None: + self.gpu_global_rank = multi_gpu_cfg["global_rank"] + self.gpu_world_size = multi_gpu_cfg["world_size"] + else: + self.gpu_global_rank = 0 + self.gpu_world_size = 1 + + if rnd_cfg: + rnd_lr = rnd_cfg.pop("learning_rate", 1e-3) + from rsl_rl.extensions import RandomNetworkDistillation + + self.rnd = RandomNetworkDistillation(device=self.device, **rnd_cfg) + self.rnd_optimizer = optim.Adam(self.rnd.predictor.parameters(), lr=rnd_lr) + else: + self.rnd = None + self.rnd_optimizer = None + + self.symmetry: dict[str, Any] | None + if symmetry_cfg is not None: + use_symmetry = symmetry_cfg["use_data_augmentation"] or symmetry_cfg["use_mirror_loss"] + if not use_symmetry: + print("Symmetry not used for learning. We will use it for logging instead.") + from rsl_rl.utils import resolve_callable + + symmetry_cfg["data_augmentation_func"] = resolve_callable( + symmetry_cfg["data_augmentation_func"] + ) + if not callable(symmetry_cfg["data_augmentation_func"]): + raise ValueError( + "Symmetry configuration exists but the function is not callable: " + f"{symmetry_cfg['data_augmentation_func']}" + ) + if actor.is_recurrent or critic.is_recurrent: + raise ValueError("Symmetry augmentation is not supported for recurrent policies.") + self.symmetry = symmetry_cfg + else: + self.symmetry = None + + self_ref = cast(Any, self) + self_ref.actor = actor.to(self.device) + self_ref.critic = critic.to(self.device) + optimizer_cls = cast(Callable[..., optim.Optimizer], resolve_optimizer(optimizer)) + self.optimizer = optimizer_cls( + self._unique_trainable_parameters(), + lr=learning_rate, + ) + self.storage = storage + self.transition = RolloutStorage.Transition() + + self.clip_param = clip_param + self.num_learning_epochs = num_learning_epochs + self.num_mini_batches = num_mini_batches + self.value_loss_coef = value_loss_coef + self.entropy_coef = entropy_coef + self.gamma = gamma + self.lam = lam + self.max_grad_norm = max_grad_norm + self.use_clipped_value_loss = use_clipped_value_loss + self.desired_kl = desired_kl + self.schedule = schedule + self.learning_rate = learning_rate + self.normalize_advantage_per_mini_batch = normalize_advantage_per_mini_batch + + def _unique_trainable_parameters(self) -> list[torch.nn.Parameter]: + params: list[torch.nn.Parameter] = [] + seen: set[int] = set() + for param in chain(self.actor.parameters(), self.critic.parameters()): + ident = id(param) + if ident in seen: + continue + seen.add(ident) + params.append(param) + return params + + @staticmethod + def construct_algorithm(obs: TensorDict, env: VecEnv, cfg: dict, device: str) -> PPO: + cfg["obs_groups"] = resolve_obs_groups(obs, cfg["obs_groups"], ["actor", "critic"]) + cfg["algorithm"] = resolve_rnd_config(cfg["algorithm"], obs, cfg["obs_groups"], env) + cfg["algorithm"] = resolve_symmetry_config(cfg["algorithm"], env) + + actor_cfg = dict(cfg["actor"]) + critic_cfg = dict(cfg["critic"]) + actor_cfg.pop("class_name", None) + critic_cfg.pop("class_name", None) + + proprio_hist_len = int(actor_cfg.pop("proprio_hist_len", obs["proprio_hist"].shape[1])) + proprio_frame_dim = int(actor_cfg.pop("proprio_frame_dim", obs["proprio_hist"].shape[-1])) + + shared_model = HoraSharedActorCritic( + obs_dim=int(obs["actor"].shape[-1]), + action_dim=int(env.num_actions), + priv_info_dim=int(obs["priv_info"].shape[-1]), + actor_hidden_dims=actor_cfg.pop("hidden_dims", (512, 256, 128)), + activation=actor_cfg.pop("activation", "elu"), + obs_normalization=bool(actor_cfg.pop("obs_normalization", False)), + distribution_cfg=actor_cfg.pop("distribution_cfg", None), + priv_info_embed_dim=int( + actor_cfg.pop("priv_info_embed_dim", obs["priv_info"].shape[-1]) + ), + priv_mlp_hidden_dims=actor_cfg.pop("priv_mlp_hidden_dims", (256, 128, 8)), + use_student_encoder=bool(actor_cfg.pop("use_student_encoder", False)), + proprio_hist_len=proprio_hist_len, + proprio_frame_dim=proprio_frame_dim, + ).to(device) + + actor = HoraActorModel( + obs, + cfg["obs_groups"], + "actor", + env.num_actions, + shared_model=shared_model, + **actor_cfg, + ).to(device) + critic = HoraCriticModel( + obs, + cfg["obs_groups"], + "critic", + 1, + shared_model=shared_model, + **critic_cfg, + ).to(device) + + storage = RolloutStorage( + "rl", env.num_envs, cfg["num_steps_per_env"], obs, [env.num_actions], device + ) + algorithm_cfg = dict(cfg["algorithm"]) + algorithm_cfg.pop("class_name", None) + return HoraPPO( + actor, critic, storage, device=device, **algorithm_cfg, multi_gpu_cfg=cfg["multi_gpu"] + ) + + def process_env_step( + self, + obs: TensorDict, + rewards: torch.Tensor, + dones: torch.Tensor, + extras: dict[str, torch.Tensor | TensorDict], + ) -> None: + self.actor.update_normalization(obs) + if self.rnd: + self.rnd.update_normalization(obs) + + self.transition.rewards = rewards.clone() + self.transition.dones = dones + + if self.rnd: + self.intrinsic_rewards = self.rnd.get_intrinsic_reward(obs) + self.transition.rewards += self.intrinsic_rewards + + timeouts = extras.get("time_outs") + timeout_bootstrap_obs = extras.get("time_out_bootstrap_obs") + if isinstance(timeouts, torch.Tensor): + timeout_mask = timeouts.to(self.device).float() + can_bootstrap = ( + timeout_bootstrap_obs is not None + and isinstance(timeout_bootstrap_obs, TensorDict) + and "priv_info" in timeout_bootstrap_obs + and torch.count_nonzero(timeout_mask) > 0 + ) + if can_bootstrap: + assert isinstance(timeout_bootstrap_obs, TensorDict) + bootstrap_obs = timeout_bootstrap_obs.to(self.device) + bootstrap_values = self.critic(bootstrap_obs).detach() + self.transition.rewards += self.gamma * torch.squeeze( + bootstrap_values * timeout_mask.unsqueeze(1), 1 + ) + else: + transition_values = self.transition.values + assert transition_values is not None + self.transition.rewards += self.gamma * torch.squeeze( + transition_values * timeout_mask.unsqueeze(1), 1 + ) + + self.storage.add_transition(self.transition) + self.transition.clear() + self.actor.reset(dones) + self.critic.reset(dones) diff --git a/src/unilab/algos/torch/hora/rsl_rl.py b/src/unilab/algos/torch/hora/rsl_rl.py new file mode 100644 index 000000000..12800d3b2 --- /dev/null +++ b/src/unilab/algos/torch/hora/rsl_rl.py @@ -0,0 +1,157 @@ +"""HORA-owned RSL-RL wrapper helpers for teacher-policy runtime.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + +import numpy as np +import torch +from tensordict import TensorDict + +from unilab.base.final_observation import resolve_terminal_observation_contract +from unilab.training.rsl_rl import RslRlVecEnvWrapper +from unilab.utils.tensor import to_numpy, to_torch + +from .observations import build_hora_obs_tensordict +from .runtime import is_hora_ppo_runtime + + +@dataclass(frozen=True) +class HoraRslRlPPORuntime: + """Resolved HORA PPO runtime consumed by the generic RSL-RL script.""" + + wrapper_cls: type[RslRlVecEnvWrapper] + + +def resolve_hora_ppo_runtime( + rl_cfg: dict[str, Any], +) -> HoraRslRlPPORuntime | None: + """Resolve HORA PPO entrypoints from an explicit runtime marker.""" + if not is_hora_ppo_runtime(rl_cfg): + return None + return HoraRslRlPPORuntime(wrapper_cls=HoraRslRlVecEnvWrapper) + + +def resolve_hora_ppo_wrapper_cls( + rl_cfg: dict[str, Any], +) -> type[RslRlVecEnvWrapper] | None: + """Return the HORA-specific PPO wrapper class when the config selects it. + + Args: + rl_cfg: Resolved algorithm config dictionary from Hydra composition. + + Returns: + ``HoraRslRlVecEnvWrapper`` when the owner config selects HORA PPO, otherwise + ``None``. + """ + runtime = resolve_hora_ppo_runtime(rl_cfg) + if runtime is None: + return None + return runtime.wrapper_cls + + +class HoraRslRlVecEnvWrapper(RslRlVecEnvWrapper): + """RSL-RL adapter that preserves HORA teacher-policy observation payloads.""" + + def _obs_to_tensordict( + self, + obs: dict[str, Any], + info: dict[str, Any] | None = None, + ) -> TensorDict: + """Convert env outputs to a HORA-aware TensorDict. + + Args: + obs: Environment observation dict following the UniLab env contract. + info: Optional env info dict containing HORA privileged payloads. + + Returns: + TensorDict preserving generic observation keys plus HORA privileged inputs. + """ + policy_obs = to_numpy(self._policy_obs(obs)) + return build_hora_obs_tensordict( + obs, + info=info, + device=self.device, + batch_size=self.num_envs, + policy_obs=policy_obs, + ) + + def step( + self, actions: torch.Tensor | np.ndarray + ) -> tuple[TensorDict, torch.Tensor, torch.Tensor, dict]: + """Step the wrapped env while keeping HORA bootstrap payloads intact. + + Args: + actions: Torch or numpy action batch with shape ``(num_envs, action_dim)``. + + Returns: + Tuple ``(obs_td, rewards, dones, infos)`` matching the RSL-RL VecEnv + contract while preserving HORA privileged observations. + """ + actions_np = to_numpy(actions) + state = self.env.step(actions_np) + rewards = to_torch(state.reward, self.device) + dones = self._resolve_done(state) + + self.episode_returns += rewards + self.episode_lengths += 1 + + infos: dict[str, torch.Tensor | TensorDict | dict[str, Any]] = {} + done_idx = torch.nonzero(dones).flatten() + if len(done_idx) > 0: + infos["time_outs"] = to_torch(state.truncated, self.device).bool() + + final_observation = self._resolve_final_observation(state) + terminal_contract = resolve_terminal_observation_contract( + next_obs_batch_size=self.num_envs, + final_observation=final_observation, + done=to_numpy(dones), + info=state.info, + truncated=to_numpy(infos["time_outs"]) if "time_outs" in infos else None, + ) + if np.any(terminal_contract.timeout_terminal_mask) and final_observation is not None: + infos["time_out_bootstrap_obs"] = self._obs_to_tensordict(final_observation) + + self.episode_returns[done_idx] = 0 + self.episode_lengths[done_idx] = 0 + + if "log" in state.info: + infos["log"] = state.info["log"] + + return ( + self._obs_to_tensordict(state.obs, state.info), + rewards, + dones, + infos, + ) + + def reset(self) -> tuple[TensorDict, dict[str, Any]]: + """Reset the wrapped env and preserve HORA privileged reset payloads. + + Args: + None. + + Returns: + Tuple ``(obs_td, info)`` where ``obs_td`` retains HORA privileged inputs. + """ + if self.env.state is None: + self.env.init_state() + + env_indices = np.arange(self.num_envs, dtype=np.int32) + obs_out, info = self.env.reset(env_indices) + self.episode_returns[:] = 0 + self.episode_lengths[:] = 0 + return self._obs_to_tensordict(obs_out, info), info + + def get_observations(self) -> TensorDict: + """Return the current HORA-aware observation TensorDict. + + Args: + None. + + Returns: + TensorDict containing the current observation batch with HORA extras. + """ + assert self.env.state is not None + return self._obs_to_tensordict(self.env.state.obs, self.env.state.info) diff --git a/src/unilab/algos/torch/hora/rsl_rl_compat.py b/src/unilab/algos/torch/hora/rsl_rl_compat.py new file mode 100644 index 000000000..670b2d7f8 --- /dev/null +++ b/src/unilab/algos/torch/hora/rsl_rl_compat.py @@ -0,0 +1,207 @@ +"""Compatibility helpers for HORA's supported RSL-RL config schemas. + +The HORA APPO code uses these helpers to normalize owner configs before +constructing grouped actor/critic modules across supported RSL-RL releases. +""" + +from __future__ import annotations + +import importlib.metadata +from copy import deepcopy +from functools import lru_cache +from typing import Any + +from packaging.version import Version + +_MLX_PPO_ONLY_KEYS = { + "adaptive_kl_beta", + "adaptive_lr_decay", + "adaptive_lr_growth", + "adaptive_lr_update_interval", + "disable_finite_checks", + "enable_compile", + "finite_check_interval", + "metrics_interval", + "target_kl_stop", + "warmup_finite_check_interval", + "warmup_metrics_interval", + "warmup_strict_iters", +} + + +@lru_cache(maxsize=1) +def get_rsl_rl_version() -> str: + """Resolve the installed RSL-RL package version. + + Args: + None. + + Returns: + Installed version string from either ``rsl-rl-lib`` or the legacy + ``rsl-rl`` package name. + """ + try: + return importlib.metadata.version("rsl-rl-lib") + except importlib.metadata.PackageNotFoundError: + try: + return importlib.metadata.version("rsl-rl") + except importlib.metadata.PackageNotFoundError as exc: + raise ImportError( + "rsl_rl is not installed. Install via: pip install rsl-rl-lib" + ) from exc + + +def is_rsl_rl_v4() -> bool: + """Check whether the active RSL-RL runtime is version 4 or newer. + + Args: + None. + + Returns: + ``True`` when the installed package version is ``>= 4.0.0``. + """ + return bool(Version(get_rsl_rl_version()) >= Version("4.0.0")) + + +def is_rsl_rl_v5() -> bool: + """Check whether the active RSL-RL runtime is version 5 or newer. + + Args: + None. + + Returns: + ``True`` when the installed package version is ``>= 5.0.0``. + """ + return bool(Version(get_rsl_rl_version()) >= Version("5.0.0")) + + +def _normalize_obs_groups_for_rsl(cfg: dict[str, Any]) -> None: + """Translate UniLab owner obs-group aliases into RSL-RL actor/critic groups. + + Args: + cfg: Mutable RSL-RL config dictionary to normalize in place. + + Returns: + None. Updates ``cfg["obs_groups"]`` directly. + """ + obs_groups_raw = cfg.get("obs_groups", {}) + obs_groups = obs_groups_raw if isinstance(obs_groups_raw, dict) else {} + + if "default" in obs_groups: + if "actor" not in obs_groups: + obs_groups["actor"] = obs_groups["default"] + if "critic" not in obs_groups: + obs_groups["critic"] = obs_groups["default"] + else: + # Keep grouped-dict specs intact for owner runtimes like HORA; the legacy + # v3 -> v4 rename only applies to flat list-based group aliases. + if isinstance(obs_groups.get("actor"), list): + obs_groups["actor"] = ["policy"] + if isinstance(obs_groups.get("critic"), list): + obs_groups["critic"] = ["policy"] + + cfg["obs_groups"] = obs_groups + + +def _convert_policy_to_actor_critic( + cfg: dict[str, Any], + *, + distribution_class_name: str, +) -> None: + """Split a legacy single ``policy`` config into ``actor`` and ``critic`` blocks. + + Args: + cfg: Mutable RSL-RL config dictionary to normalize in place. + distribution_class_name: Distribution class name expected by the target + RSL-RL runtime. + + Returns: + None. Updates ``cfg`` directly when a legacy ``policy`` block is present. + """ + empirical_normalization = bool(cfg.pop("empirical_normalization", False)) + cfg.pop("runner_class_name", None) + + if "policy" not in cfg or "actor" in cfg or "critic" in cfg: + return + + policy = cfg.pop("policy") + if not isinstance(policy, dict): + return + + cfg["actor"] = { + "class_name": "MLPModel", + "hidden_dims": policy.get("actor_hidden_dims", [256, 256, 256]), + "activation": policy.get("activation", "elu"), + "obs_normalization": empirical_normalization, + "distribution_cfg": { + "class_name": distribution_class_name, + "init_std": policy.get("init_noise_std", 1.0), + "std_type": policy.get("noise_std_type", "scalar"), + }, + } + cfg["critic"] = { + "class_name": "MLPModel", + "hidden_dims": policy.get("critic_hidden_dims", [256, 256, 256]), + "activation": policy.get("activation", "elu"), + "obs_normalization": empirical_normalization, + } + + +def _normalize_algorithm_cfg(cfg: dict[str, Any]) -> None: + """Remove owner-only keys that current RSL-RL releases do not accept. + + Args: + cfg: Mutable RSL-RL config dictionary to normalize in place. + + Returns: + None. Updates ``cfg["algorithm"]`` directly when present. + """ + algorithm_cfg = cfg.get("algorithm") + if not isinstance(algorithm_cfg, dict): + return + + algorithm_cfg.setdefault("rnd_cfg", None) + for key in _MLX_PPO_ONLY_KEYS: + algorithm_cfg.pop(key, None) + + +def convert_config_v3_to_v4(cfg: dict[str, Any]) -> dict[str, Any]: + """Convert a legacy UniLab PPO/APPO config into the RSL-RL v4 schema. + + Args: + cfg: Resolved owner config dictionary before RSL-RL construction. + + Returns: + Deep-copied config dictionary aligned with the RSL-RL v4 actor/critic + schema and obs-group naming. + """ + converted = deepcopy(cfg) + _convert_policy_to_actor_critic( + converted, + distribution_class_name="rsl_rl.modules.distribution.GaussianDistribution", + ) + _normalize_algorithm_cfg(converted) + _normalize_obs_groups_for_rsl(converted) + if "multi_gpu" not in converted: + converted["multi_gpu"] = None + return converted + + +def convert_config_v5(cfg: dict[str, Any]) -> dict[str, Any]: + """Convert a legacy UniLab PPO/APPO config into the RSL-RL v5 schema. + + Args: + cfg: Resolved owner config dictionary before RSL-RL construction. + + Returns: + Deep-copied config dictionary aligned with the RSL-RL v5 actor/critic + schema and obs-group naming. + """ + converted = deepcopy(cfg) + _convert_policy_to_actor_critic( + converted, + distribution_class_name="GaussianDistribution", + ) + _normalize_algorithm_cfg(converted) + _normalize_obs_groups_for_rsl(converted) + return converted diff --git a/src/unilab/algos/torch/hora/runtime.py b/src/unilab/algos/torch/hora/runtime.py new file mode 100644 index 000000000..ab1d27032 --- /dev/null +++ b/src/unilab/algos/torch/hora/runtime.py @@ -0,0 +1,47 @@ +"""Config-driven runtime selection helpers for HORA teacher-policy RL.""" + +from __future__ import annotations + +from typing import Any + +HORA_APPO_RUNTIME_IMPL = "hora_appo" +HORA_PPO_RUNTIME_IMPL = "hora_ppo" + + +def resolve_hora_runtime_impl(rl_cfg: dict[str, Any]) -> str | None: + """Return the explicit HORA runtime marker from a resolved algo config. + + Args: + rl_cfg: Resolved algorithm config dictionary from Hydra composition. + + Returns: + Runtime marker string when the owner YAML selected one, otherwise ``None``. + """ + runtime_impl = rl_cfg.get("runtime_impl") + if runtime_impl in (None, ""): + return None + return str(runtime_impl) + + +def is_hora_appo_runtime(rl_cfg: dict[str, Any]) -> bool: + """Check whether the resolved algo config selects the HORA APPO runtime. + + Args: + rl_cfg: Resolved algorithm config dictionary from Hydra composition. + + Returns: + ``True`` when the config explicitly selects the HORA APPO runtime. + """ + return resolve_hora_runtime_impl(rl_cfg) == HORA_APPO_RUNTIME_IMPL + + +def is_hora_ppo_runtime(rl_cfg: dict[str, Any]) -> bool: + """Check whether the resolved algo config selects the HORA PPO runtime. + + Args: + rl_cfg: Resolved algorithm config dictionary from Hydra composition. + + Returns: + ``True`` when the config explicitly selects the HORA PPO runtime. + """ + return resolve_hora_runtime_impl(rl_cfg) == HORA_PPO_RUNTIME_IMPL diff --git a/src/unilab/algos/torch/rsl_rl_runtime.py b/src/unilab/algos/torch/rsl_rl_runtime.py new file mode 100644 index 000000000..8afa405d6 --- /dev/null +++ b/src/unilab/algos/torch/rsl_rl_runtime.py @@ -0,0 +1,49 @@ +"""Runtime resolution helpers for RSL-RL PPO script assembly.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + +from unilab.training.rsl_rl import RslRlVecEnvWrapper + + +@dataclass(frozen=True) +class RslRlPPORuntime: + """Resolved PPO runtime consumed by the generic RSL-RL entrypoint.""" + + wrapper_cls: type[RslRlVecEnvWrapper] + + +def resolve_rsl_rl_ppo_runtime( + rl_cfg: dict[str, Any], + *, + default_wrapper_cls: type[RslRlVecEnvWrapper], +) -> RslRlPPORuntime: + """Resolve the PPO runtime bundle from owner config.""" + runtime_resolver = rl_cfg.get("runtime_resolver") + if runtime_resolver in (None, ""): + runtime_impl = rl_cfg.get("runtime_impl") + if runtime_impl not in (None, ""): + raise ValueError( + "PPO owner config selected " + f"runtime_impl={runtime_impl!r} but did not define algo.runtime_resolver." + ) + return RslRlPPORuntime(wrapper_cls=default_wrapper_cls) + + from rsl_rl.utils import resolve_callable + + resolver = resolve_callable(str(runtime_resolver)) + runtime = resolver(rl_cfg) + if runtime is None: + raise ValueError( + f"PPO runtime resolver {runtime_resolver!r} returned None for rl_cfg runtime selection." + ) + + wrapper_cls = getattr(runtime, "wrapper_cls", None) + if wrapper_cls is None: + raise TypeError( + f"PPO runtime resolver {runtime_resolver!r} must return an object with " + "'wrapper_cls' attribute." + ) + return RslRlPPORuntime(wrapper_cls=wrapper_cls) diff --git a/src/unilab/assets/robots/allegro_hand/scene.xml b/src/unilab/assets/robots/allegro_hand/scene.xml index c2ce2d1ff..d0b19869f 100644 --- a/src/unilab/assets/robots/allegro_hand/scene.xml +++ b/src/unilab/assets/robots/allegro_hand/scene.xml @@ -1,4 +1,4 @@ - + diff --git a/src/unilab/base/backend/base.py b/src/unilab/base/backend/base.py index 422b8334b..89ae7b218 100644 --- a/src/unilab/base/backend/base.py +++ b/src/unilab/base/backend/base.py @@ -73,6 +73,10 @@ def get_keyframe_qpos(self, name: str) -> np.ndarray: (nq,) 数组 """ + def get_default_qpos(self) -> np.ndarray: + """Return the backend/model default qpos through a stable contract.""" + raise NotImplementedError(f"{self.__class__.__name__} does not expose default qpos") + @abc.abstractmethod def get_init_qvel(self) -> np.ndarray: """获取零初始化的 qvel 向量,维度与 set_state 期望一致 @@ -95,6 +99,50 @@ def get_body_ids(self, names: Sequence[str]) -> np.ndarray: ValueError: 若名称未找到 """ + def get_body_id(self, name: str) -> int: + """Resolve one body/link name through the backend contract.""" + return int(self.get_body_ids([name])[0]) + + def get_geom_id(self, name: str) -> int: + """Resolve one geom name through the backend contract.""" + raise NotImplementedError(f"{self.__class__.__name__} does not expose geom ids") + + def get_geom_size(self, name: str) -> np.ndarray: + """Return one geom size vector through the backend contract.""" + raise NotImplementedError(f"{self.__class__.__name__} does not expose geom sizes") + + def get_body_subtree_ids(self, root_body_id: int) -> np.ndarray: + """Return body ids in the subtree rooted at ``root_body_id``.""" + raise NotImplementedError(f"{self.__class__.__name__} does not expose body subtree ids") + + def get_geom_names(self) -> tuple[str, ...]: + """Return backend geom names in backend id order.""" + raise NotImplementedError(f"{self.__class__.__name__} does not expose geom names") + + def get_geom_body_ids(self) -> np.ndarray: + """Return the owning body id for each geom.""" + raise NotImplementedError(f"{self.__class__.__name__} does not expose geom body ids") + + def get_geom_contact_masks(self) -> tuple[np.ndarray, np.ndarray]: + """Return per-geom contact type and affinity masks.""" + raise NotImplementedError(f"{self.__class__.__name__} does not expose geom contact masks") + + def get_geom_friction(self) -> np.ndarray: + """Return the backend geom-friction table.""" + raise NotImplementedError(f"{self.__class__.__name__} does not expose geom friction") + + def get_gravity(self) -> np.ndarray: + """Return the backend gravity vector.""" + raise NotImplementedError(f"{self.__class__.__name__} does not expose gravity") + + def get_body_mass(self) -> np.ndarray: + """Return the backend body-mass table.""" + raise NotImplementedError(f"{self.__class__.__name__} does not expose body mass") + + def get_body_ipos(self) -> np.ndarray: + """Return the backend body inertial-position table.""" + raise NotImplementedError(f"{self.__class__.__name__} does not expose body ipos") + def get_motion_body_ids(self, names: Sequence[str]) -> np.ndarray: """Resolve MuJoCo-style body IDs used by motion datasets.""" from unilab.base.backend.xml import get_named_body_ids @@ -161,6 +209,42 @@ def materialize(self) -> None: def apply_interval_randomization(self, plan: IntervalRandomizationPlan) -> None: """Apply a scheduled interval randomization plan.""" + def apply_body_linear_velocity_delta( + self, + body_ids: np.ndarray, + velocity_delta: np.ndarray, + ) -> None: + """Apply a world-frame linear-velocity delta to specific bodies. + + Args: + body_ids: Body ids whose linear velocities should be perturbed. + velocity_delta: Velocity delta with shape ``(num_envs, len(body_ids), 3)``. + + Returns: + None. Backends that support this mutate their pending simulation state. + """ + raise NotImplementedError( + f"{self.__class__.__name__} does not support interval body velocity perturbation" + ) + + def apply_body_force( + self, + body_ids: np.ndarray, + force: np.ndarray, + ) -> None: + """Apply a world-frame force to specific bodies for the upcoming step. + + Args: + body_ids: Body ids whose external forces should be perturbed. + force: Force values with shape ``(num_envs, len(body_ids), 3)``. + + Returns: + None. Backends that support this mutate their pending simulation state. + """ + raise NotImplementedError( + f"{self.__class__.__name__} does not support interval body force perturbation" + ) + def get_play_capabilities(self) -> BackendPlayCapabilities: """Return backend-native play/render capabilities.""" return BackendPlayCapabilities() diff --git a/src/unilab/base/backend/motrix_backend.py b/src/unilab/base/backend/motrix_backend.py index 655e7ab19..a24263a1c 100644 --- a/src/unilab/base/backend/motrix_backend.py +++ b/src/unilab/base/backend/motrix_backend.py @@ -187,6 +187,9 @@ def get_keyframe_qpos(self, name: str) -> np.ndarray: return np.array(self._model.keyframes[0].dof_pos, dtype=self._np_dtype) return np.array(self._model.compute_init_dof_pos(), dtype=self._np_dtype) + def get_default_qpos(self) -> np.ndarray: + return np.array(self._model.compute_init_dof_pos(), dtype=self._np_dtype) + def get_init_qvel(self) -> np.ndarray: return np.zeros((self._model.num_dof_vel,), dtype=self._np_dtype) @@ -263,6 +266,7 @@ def get_dr_capabilities(self) -> DomainRandomizationCapabilities: {RESET_TERM_BASE_MASS, RESET_TERM_BASE_COM, RESET_TERM_KP, RESET_TERM_KD} ), supports_interval_push=True, + supports_interval_body_velocity_delta=False, ) def apply_interval_randomization(self, plan: IntervalRandomizationPlan) -> None: diff --git a/src/unilab/base/backend/mujoco_backend.py b/src/unilab/base/backend/mujoco_backend.py index a8bac0b23..649c9298b 100644 --- a/src/unilab/base/backend/mujoco_backend.py +++ b/src/unilab/base/backend/mujoco_backend.py @@ -3,7 +3,7 @@ import time from collections.abc import Sequence from concurrent.futures import ProcessPoolExecutor -from multiprocessing import cpu_count, get_context +from multiprocessing import cpu_count, current_process, get_context from typing import Optional, cast import mujoco @@ -14,7 +14,10 @@ RESET_TERM_BASE_COM, RESET_TERM_BASE_MASS, RESET_TERM_BODY_INERTIA, + RESET_TERM_BODY_IPOS, RESET_TERM_BODY_IQUAT, + RESET_TERM_BODY_MASS, + RESET_TERM_GEOM_FRICTION, RESET_TERM_GRAVITY, RESET_TERM_KD, RESET_TERM_KP, @@ -284,6 +287,19 @@ def _resolve_push_body_force_slice(self, body_id: int) -> slice: start = 6 * body_id return slice(start, start + 3) + def _sample_push_force(self, force_range: Sequence[float] | np.ndarray) -> np.ndarray: + """Sample one world-frame push force vector per environment. + + Args: + force_range: Per-axis push-force magnitude range. + + Returns: + Array with shape ``(num_envs, 3)`` containing sampled forces. + """ + ex_force = np.random.uniform(-1.0, 1.0, size=(self._num_envs, 3)) + ex_force *= np.asarray(force_range, dtype=np.float64) + return ex_force.astype(np.float64, copy=False) + def _compile_model_variants( self, variant_specs: Sequence[ModelVariantSpec], @@ -291,30 +307,33 @@ def _compile_model_variants( variants = tuple(variant_specs) if not variants: return tuple() - if len(variants) == 1: - mjb_paths = _compile_model_variant_chunk_to_mjb( - model_file=self._model_file, - add_body_sensors=self.add_body_sensors, - base_name=self._base_name, - sim_dt=self._sim_dt, - iterations=self._iterations, - position_actuator_gains=self._position_actuator_gains, - variants=variants, - ) + + def _load_compiled_models_and_cleanup(paths: Sequence[str]) -> tuple[mujoco.MjModel, ...]: try: - models = tuple(mujoco.MjModel.from_binary_path(path) for path in mjb_paths) + return tuple(mujoco.MjModel.from_binary_path(path) for path in paths) finally: - for path in mjb_paths: + for path in paths: if os.path.exists(path): os.remove(path) - for path in mjb_paths: + for path in paths: parent = os.path.dirname(path) if parent and os.path.isdir(parent): try: os.rmdir(parent) except OSError: pass - return models + + if len(variants) == 1 or current_process().daemon: + mjb_paths = _compile_model_variant_chunk_to_mjb( + model_file=self._model_file, + add_body_sensors=self.add_body_sensors, + base_name=self._base_name, + sim_dt=self._sim_dt, + iterations=self._iterations, + position_actuator_gains=self._position_actuator_gains, + variants=variants, + ) + return _load_compiled_models_and_cleanup(mjb_paths) max_workers = min(len(variants), max(1, cpu_count())) chunk_size = max(1, (len(variants) + max_workers - 1) // max_workers) @@ -354,20 +373,7 @@ def _compile_model_variants( for chunk in chunks ] flat_paths = [path for paths in mjb_paths_nested for path in paths] - try: - models = tuple(mujoco.MjModel.from_binary_path(path) for path in flat_paths) - finally: - for path in flat_paths: - if os.path.exists(path): - os.remove(path) - for path in flat_paths: - parent = os.path.dirname(path) - if parent and os.path.isdir(parent): - try: - os.rmdir(parent) - except OSError: - pass - return models + return _load_compiled_models_and_cleanup(flat_paths) def _current_model_sequence(self) -> mujoco.MjModel | list[mujoco.MjModel]: if len(self._model_variants) == 1 and np.all(self._model_assignments == 0): @@ -450,6 +456,9 @@ def get_keyframe_qpos(self, name: str) -> np.ndarray: raise ValueError(f"Keyframe '{name}' not found in MuJoCo model") return np.array(self._model.key_qpos[key_id].copy(), dtype=self._np_dtype) + def get_default_qpos(self) -> np.ndarray: + return np.asarray(self._model.qpos0, dtype=np.float64).copy() + def get_init_qvel(self) -> np.ndarray: return np.zeros((self.nv,), dtype=self._np_dtype) @@ -462,6 +471,54 @@ def get_body_ids(self, names: "Sequence[str]") -> np.ndarray: ids.append(bid) return np.array(ids, dtype=np.int32) + def get_geom_id(self, name: str) -> int: + geom_id = mujoco.mj_name2id(self._model, mujoco.mjtObj.mjOBJ_GEOM, name) + if geom_id < 0: + raise ValueError(f"Geom '{name}' not found in MuJoCo model") + return int(geom_id) + + def get_geom_size(self, name: str) -> np.ndarray: + return np.asarray(self._model.geom_size[self.get_geom_id(name)], dtype=np.float64).copy() + + def get_body_subtree_ids(self, root_body_id: int) -> np.ndarray: + subtree_ids = {int(root_body_id)} + changed = True + while changed: + changed = False + for body_id in range(self._model.nbody): + parent_id = int(self._model.body_parentid[body_id]) + if body_id not in subtree_ids and parent_id in subtree_ids: + subtree_ids.add(body_id) + changed = True + return np.asarray(sorted(subtree_ids), dtype=np.int32) + + def get_geom_names(self) -> tuple[str, ...]: + return tuple( + mujoco.mj_id2name(self._model, mujoco.mjtObj.mjOBJ_GEOM, geom_id) or "" + for geom_id in range(self._model.ngeom) + ) + + def get_geom_body_ids(self) -> np.ndarray: + return np.asarray(self._model.geom_bodyid, dtype=np.int32).copy() + + def get_geom_contact_masks(self) -> tuple[np.ndarray, np.ndarray]: + return ( + np.asarray(self._model.geom_contype, dtype=np.int32).copy(), + np.asarray(self._model.geom_conaffinity, dtype=np.int32).copy(), + ) + + def get_geom_friction(self) -> np.ndarray: + return np.asarray(self._model.geom_friction, dtype=np.float64).copy() + + def get_gravity(self) -> np.ndarray: + return np.asarray(self._model.opt.gravity, dtype=np.float64).copy() + + def get_body_mass(self) -> np.ndarray: + return np.asarray(self._model.body_mass, dtype=np.float64).copy() + + def get_body_ipos(self) -> np.ndarray: + return np.asarray(self._model.body_ipos, dtype=np.float64).copy() + def get_motion_body_ids(self, names: Sequence[str]) -> np.ndarray: return self.get_body_ids(names) @@ -545,11 +602,15 @@ def get_dr_capabilities(self) -> DomainRandomizationCapabilities: RESET_TERM_GRAVITY, RESET_TERM_BODY_IQUAT, RESET_TERM_BODY_INERTIA, + RESET_TERM_BODY_IPOS, + RESET_TERM_BODY_MASS, + RESET_TERM_GEOM_FRICTION, RESET_TERM_KP, RESET_TERM_KD, } ), supports_interval_push=self._push_body_id >= 0, + supports_interval_body_force=True, ) def apply_init_randomization(self, plan: InitRandomizationPlan) -> None: @@ -567,15 +628,49 @@ def materialize(self) -> None: self._pool = self._build_pool() def apply_interval_randomization(self, plan: IntervalRandomizationPlan) -> None: - if plan.push_perturbation_limit is None: + if plan.is_empty(): return - self.push_robots(plan.push_perturbation_limit) + self._pending_xfrc_applied.fill(0.0) + if plan.push_perturbation_limit is not None: + self.push_robots(plan.push_perturbation_limit) + if plan.body_force is not None: + if plan.body_ids is None: + raise ValueError("Interval body-force perturbation requires body_ids") + self.apply_body_force(plan.body_ids, plan.body_force) + if plan.body_linear_velocity_delta is not None: + if plan.body_ids is None: + raise ValueError("Interval body-velocity perturbation requires body_ids") + self.apply_body_linear_velocity_delta(plan.body_ids, plan.body_linear_velocity_delta) def push_robots(self, force_range: Sequence[float] | np.ndarray) -> None: - ex_force = np.random.uniform(-1.0, 1.0, size=(self._num_envs, 3)) - ex_force *= force_range self._pending_xfrc_applied.fill(0.0) - self._pending_xfrc_applied[:, self._push_body_force_slice] = ex_force + self._pending_xfrc_applied[:, self._push_body_force_slice] = self._sample_push_force( + force_range + ) + + def apply_body_force( + self, + body_ids: np.ndarray, + force: np.ndarray, + ) -> None: + """Accumulate one external world-frame force vector per target body. + + Args: + body_ids: Body ids to perturb. + force: Force tensor with shape ``(num_envs, len(body_ids), 3)``. + + Returns: + None. The force is staged in ``xfrc_applied`` for the next step. + """ + body_ids_np = np.asarray(body_ids, dtype=np.int32).reshape(-1) + force_np = np.asarray(force, dtype=np.float64) + expected_shape = (self._num_envs, body_ids_np.size, 3) + if force_np.shape != expected_shape: + raise ValueError(f"body force must have shape {expected_shape}, got {force_np.shape}") + for body_offset, body_id in enumerate(body_ids_np): + self._pending_xfrc_applied[:, self._resolve_push_body_force_slice(int(body_id))] += ( + force_np[:, body_offset, :] + ) def get_play_capabilities(self) -> BackendPlayCapabilities: return BackendPlayCapabilities(supports_physics_state_playback=True) @@ -703,16 +798,38 @@ def _translate_reset_randomization( raise ValueError(f"Body '{self._base_name}' not found in MuJoCo model") translated: dict[str, np.ndarray] = {} + body_mass = None + if randomization.body_mass is not None: + body_mass = self._coerce_reset_field( + randomization.body_mass, + name="body_mass", + num_reset=num_reset, + shaped_tail=(self._model.nbody,), + ) if randomization.base_mass_delta is not None: - body_mass = np.broadcast_to(self._base_body_mass, (num_reset, self._model.nbody)).copy() + if body_mass is None: + body_mass = np.broadcast_to( + self._base_body_mass, (num_reset, self._model.nbody) + ).copy() body_mass[:, self._base_body_id] += np.asarray(randomization.base_mass_delta) + if body_mass is not None: translated["body_mass"] = body_mass + body_ipos = None + if randomization.body_ipos is not None: + body_ipos = self._coerce_reset_field( + randomization.body_ipos, + name="body_ipos", + num_reset=num_reset, + shaped_tail=(self._model.nbody, 3), + ) if randomization.base_com_offset is not None: - body_ipos = np.broadcast_to( - self._base_body_ipos, (num_reset, self._model.nbody, 3) - ).copy() + if body_ipos is None: + body_ipos = np.broadcast_to( + self._base_body_ipos, (num_reset, self._model.nbody, 3) + ).copy() body_ipos[:, self._base_body_id, :] += np.asarray(randomization.base_com_offset) + if body_ipos is not None: translated["body_ipos"] = body_ipos.reshape(num_reset, -1) if randomization.gravity is not None: @@ -739,6 +856,14 @@ def _translate_reset_randomization( shaped_tail=(self._model.nbody, 3), ) + if randomization.geom_friction is not None: + translated["geom_friction"] = self._coerce_reset_field( + randomization.geom_friction, + name="geom_friction", + num_reset=num_reset, + shaped_tail=(self._model.ngeom, 3), + ) + if randomization.kp is not None: translated["kp"] = self._coerce_reset_field( randomization.kp, diff --git a/src/unilab/dr/manager.py b/src/unilab/dr/manager.py index 78a41fcd9..4d689873d 100644 --- a/src/unilab/dr/manager.py +++ b/src/unilab/dr/manager.py @@ -53,6 +53,17 @@ def apply_interval_randomization_if_due(self, step_counter: int) -> None: raise NotImplementedError( f"{self._env._backend.backend_type} backend does not support interval push" ) + if ( + plan.body_linear_velocity_delta is not None + and not self._capabilities.supports_interval_body_velocity_delta + ): + raise NotImplementedError( + f"{self._env._backend.backend_type} backend does not support interval body velocity perturbation" + ) + if plan.body_force is not None and not self._capabilities.supports_interval_body_force: + raise NotImplementedError( + f"{self._env._backend.backend_type} backend does not support interval body force perturbation" + ) self._env._backend.apply_interval_randomization(plan) def _log_unsupported_reset_terms(self, unsupported: frozenset[str]) -> None: diff --git a/src/unilab/dr/types.py b/src/unilab/dr/types.py index dfd670fcb..6be7ea641 100644 --- a/src/unilab/dr/types.py +++ b/src/unilab/dr/types.py @@ -11,6 +11,9 @@ RESET_TERM_GRAVITY = "gravity" RESET_TERM_BODY_IQUAT = "body_iquat" RESET_TERM_BODY_INERTIA = "body_inertia" +RESET_TERM_BODY_IPOS = "body_ipos" +RESET_TERM_BODY_MASS = "body_mass" +RESET_TERM_GEOM_FRICTION = "geom_friction" RESET_TERM_KP = "kp" RESET_TERM_KD = "kd" @@ -33,6 +36,8 @@ def is_empty(self) -> bool: class DomainRandomizationCapabilities: supported_reset_terms: frozenset[str] = field(default_factory=frozenset) supports_interval_push: bool = False + supports_interval_body_velocity_delta: bool = False + supports_interval_body_force: bool = False def supports_reset_term(self, term: str) -> bool: return term in self.supported_reset_terms @@ -61,6 +66,17 @@ def filter_reset_payload( body_inertia=( payload.body_inertia if self.supports_reset_term(RESET_TERM_BODY_INERTIA) else None ), + body_ipos=( + payload.body_ipos if self.supports_reset_term(RESET_TERM_BODY_IPOS) else None + ), + body_mass=( + payload.body_mass if self.supports_reset_term(RESET_TERM_BODY_MASS) else None + ), + geom_friction=( + payload.geom_friction + if self.supports_reset_term(RESET_TERM_GEOM_FRICTION) + else None + ), kp=payload.kp if self.supports_reset_term(RESET_TERM_KP) else None, kd=payload.kd if self.supports_reset_term(RESET_TERM_KD) else None, ) @@ -74,6 +90,9 @@ class ResetRandomizationPayload: gravity: np.ndarray | None = None body_iquat: np.ndarray | None = None body_inertia: np.ndarray | None = None + body_ipos: np.ndarray | None = None + body_mass: np.ndarray | None = None + geom_friction: np.ndarray | None = None kp: np.ndarray | None = None kd: np.ndarray | None = None @@ -89,6 +108,12 @@ def requested_terms(self) -> frozenset[str]: terms.add(RESET_TERM_BODY_IQUAT) if self.body_inertia is not None: terms.add(RESET_TERM_BODY_INERTIA) + if self.body_ipos is not None: + terms.add(RESET_TERM_BODY_IPOS) + if self.body_mass is not None: + terms.add(RESET_TERM_BODY_MASS) + if self.geom_friction is not None: + terms.add(RESET_TERM_GEOM_FRICTION) if self.kp is not None: terms.add(RESET_TERM_KP) if self.kd is not None: @@ -102,9 +127,16 @@ def is_empty(self) -> bool: @dataclass class IntervalRandomizationPlan: push_perturbation_limit: Sequence[float] | np.ndarray | None = None + body_ids: np.ndarray | None = None + body_linear_velocity_delta: np.ndarray | None = None + body_force: np.ndarray | None = None def is_empty(self) -> bool: - return self.push_perturbation_limit is None + return ( + self.push_perturbation_limit is None + and self.body_linear_velocity_delta is None + and self.body_force is None + ) @dataclass diff --git a/src/unilab/envs/manipulation/__init__.py b/src/unilab/envs/manipulation/__init__.py index 5def847e4..02b0f854f 100644 --- a/src/unilab/envs/manipulation/__init__.py +++ b/src/unilab/envs/manipulation/__init__.py @@ -1,6 +1,6 @@ """Manipulation env registry bootstrap contract.""" __unilab_registry_modules__ = ( - "unilab.envs.manipulation.inhand_rot_allegro", + "unilab.envs.manipulation.allegro_inhand", "unilab.envs.manipulation.sharpa_inhand", ) diff --git a/src/unilab/envs/manipulation/inhand_rot_allegro/__init__.py b/src/unilab/envs/manipulation/allegro_inhand/__init__.py similarity index 100% rename from src/unilab/envs/manipulation/inhand_rot_allegro/__init__.py rename to src/unilab/envs/manipulation/allegro_inhand/__init__.py diff --git a/src/unilab/envs/manipulation/inhand_rot_allegro/base.py b/src/unilab/envs/manipulation/allegro_inhand/base.py similarity index 100% rename from src/unilab/envs/manipulation/inhand_rot_allegro/base.py rename to src/unilab/envs/manipulation/allegro_inhand/base.py diff --git a/src/unilab/envs/manipulation/inhand_rot_allegro/grasp_gen.py b/src/unilab/envs/manipulation/allegro_inhand/grasp_gen.py similarity index 98% rename from src/unilab/envs/manipulation/inhand_rot_allegro/grasp_gen.py rename to src/unilab/envs/manipulation/allegro_inhand/grasp_gen.py index a2f8f89ca..d94acad89 100644 --- a/src/unilab/envs/manipulation/inhand_rot_allegro/grasp_gen.py +++ b/src/unilab/envs/manipulation/allegro_inhand/grasp_gen.py @@ -9,11 +9,8 @@ from unilab.base import registry from unilab.base.np_env import NpEnvState -from unilab.envs.manipulation.inhand_rot_allegro.rotation import ( - AllegroRotationPPO, - AllegroRotationPPOCfg, - RewardConfigPPO, -) + +from .rotation import AllegroRotationPPO, AllegroRotationPPOCfg, RewardConfigPPO @registry.envcfg("AllegroInhandRotationGrasp") diff --git a/src/unilab/envs/manipulation/inhand_rot_allegro/rotation.py b/src/unilab/envs/manipulation/allegro_inhand/rotation.py similarity index 99% rename from src/unilab/envs/manipulation/inhand_rot_allegro/rotation.py rename to src/unilab/envs/manipulation/allegro_inhand/rotation.py index 376ef078e..e465d3b68 100644 --- a/src/unilab/envs/manipulation/inhand_rot_allegro/rotation.py +++ b/src/unilab/envs/manipulation/allegro_inhand/rotation.py @@ -27,7 +27,8 @@ ) from unilab.dtype_config import get_global_dtype from unilab.envs.common.rotation import np_quat_conjugate, np_quat_mul, np_quat_to_axis_angle -from unilab.envs.manipulation.inhand_rot_allegro.base import AllegroBaseCfg, AllegroBaseEnv + +from .base import AllegroBaseCfg, AllegroBaseEnv def normalize_rotation_axis(rotation_axis: tuple[float, float, float]) -> np.ndarray: diff --git a/src/unilab/envs/manipulation/sharpa_inhand/base.py b/src/unilab/envs/manipulation/sharpa_inhand/base.py index b0bb49ccd..cb672646f 100644 --- a/src/unilab/envs/manipulation/sharpa_inhand/base.py +++ b/src/unilab/envs/manipulation/sharpa_inhand/base.py @@ -78,8 +78,14 @@ @dataclass class SharpaControlConfig: action_scale: float = 1.0 / 24.0 + # MuJoCo Sharpa loads PD defaults from XML actuator gains. These fields remain + # fallback defaults for backends that cannot expose actuator gains yet. p_gain: float = 1.0 d_gain: float = 0.1 + # MuJoCo Sharpa currently uses position actuators, so torque-control mode + # stays declared here for owner-config structure but is rejected at runtime. + torque_control: bool = False + dof_limits_scale: float = 0.9 @dataclass @@ -87,8 +93,26 @@ class SharpaSensorConfig: tactile_force_sensor_names: list[str] = field(default_factory=list) +@dataclass +class SharpaObservationConfig: + observation_mode: str = "separated" + enable_tactile: bool = True + binary_contact: bool = False + enable_contact_pos: bool = False + contact_smooth: float = 0.5 + contact_threshold: float = 0.05 + tactile_force_clip_max: float = 4.0 + + +@dataclass +class SharpaPrivilegedInfoConfig: + include_friction_scale: bool = True + include_gravity_direction: bool = False + + @dataclass class SharpaDomainRandConfig: + scale_list: list[float] = field(default_factory=lambda: [0.5]) randomize_base_mass: bool = False added_mass_range: list[float] = field(default_factory=lambda: [0.0, 0.0]) random_com: bool = False @@ -97,6 +121,32 @@ class SharpaDomainRandConfig: gravity_range: list[list[float]] = field( default_factory=lambda: [[0.0, 0.0, -9.81], [0.0, 0.0, -9.81]] ) + randomize_gravity_direction: bool = False + gravity_direction_magnitude: float = 9.81 + randomize_pd_gains: bool = True + randomize_p_gain_scale_lower: float = 0.5 + randomize_p_gain_scale_upper: float = 2.0 + randomize_d_gain_scale_lower: float = 0.5 + randomize_d_gain_scale_upper: float = 2.0 + randomize_friction: bool = True + randomize_friction_scale_lower: float = 0.5 + randomize_friction_scale_upper: float = 2.0 + elastomer_base_friction: float = 1.6 + metal_base_friction: float = 0.2 + object_base_friction: float = 1.0 + randomize_com: bool = True + randomize_com_lower: float = -0.01 + randomize_com_upper: float = 0.01 + randomize_mass: bool = True + randomize_mass_lower: float = 0.01 + randomize_mass_upper: float = 0.25 + force_scale: float = 2.0 + random_force_prob_scalar: float = 0.25 + force_decay: float = 0.9 + force_decay_interval: float = 0.08 + joint_noise_scale: float = 0.02 + contact_latency: float = 0.005 + contact_sensor_noise: float = 0.01 push_body_name: str | None = None @@ -111,14 +161,9 @@ class SharpaInhandBaseCfg(EnvCfg): observation_space: int = 192 prop_hist_len: int = 30 critic_info_dim: int = 8 - # "separate": keep critic-only info in its own obs group. - # "merged": append critic info into the main "obs" vector. - critic_obs_mode: str = "separate" clip_obs: float = 5.0 clip_actions: float = 1.0 - # NOTE: Sharpa MuJoCo XML uses position actuators; true torque-control mode is not implemented. - torque_control: bool = False num_hand_dofs: int = 22 frame_obs_dim: int = 64 @@ -137,98 +182,88 @@ class SharpaInhandBaseCfg(EnvCfg): control_config: SharpaControlConfig = field(default_factory=SharpaControlConfig) sensor: SharpaSensorConfig = field(default_factory=SharpaSensorConfig) # type: ignore[assignment] + obs: SharpaObservationConfig = field(default_factory=SharpaObservationConfig) + priv_info: SharpaPrivilegedInfoConfig = field(default_factory=SharpaPrivilegedInfoConfig) domain_rand: SharpaDomainRandConfig = field(default_factory=SharpaDomainRandConfig) reset_height_lower: float = 0.59906 reset_height_upper: float = 0.63906 reset_angle_diff: float = 45.0 / 180.0 * np.pi - reset_random_quat: bool = False rot_axis: tuple[float, float, float] = (0.0, 0.0, 1.0) grasp_cache_path: str = "cache/sharpa_grasp_linspace" - - joint_noise_scale: float = 0.02 - - enable_tactile: bool = True - binary_contact: bool = False - enable_contact_pos: bool = False disable_tactile_ids: list[int] = field(default_factory=list) - contact_smooth: float = 0.5 - contact_threshold: float = 0.05 - contact_latency: float = 0.005 - contact_sensor_noise: float = 0.01 - - dof_limits_scale: float = 0.9 + # Match the reference Sharpa object-position reward/privileged-info anchor + # by using the fixed XML/default object pose instead of the sampled grasp reset. + use_default_object_pose_for_object_pos_anchor: bool = False - scale_range: list[float] = field(default_factory=lambda: [0.5, 0.5, 1.0]) - - randomize_pd_gains: bool = True - randomize_p_gain_scale_lower: float = 0.5 - randomize_p_gain_scale_upper: float = 2.0 - randomize_d_gain_scale_lower: float = 0.5 - randomize_d_gain_scale_upper: float = 2.0 + debug_show_axes: bool = False - randomize_friction: bool = True - randomize_friction_scale_lower: float = 0.5 - randomize_friction_scale_upper: float = 2.0 - elastomer_base_friction: float = 0.8 - metal_base_friction: float = 0.1 - object_base_friction: float = 0.5 - randomize_com: bool = True - randomize_com_lower: float = -0.01 - randomize_com_upper: float = 0.01 +def format_scale_tag(scale_value: float) -> str: + """Convert one object scale into a stable cache filename tag. - randomize_mass: bool = True - randomize_mass_lower: float = 0.01 - randomize_mass_upper: float = 0.25 - - force_scale: float = 2.0 - random_force_prob_scalar: float = 0.25 - force_decay: float = 0.9 - force_decay_interval: float = 0.08 + Args: + scale_value: Single object scale value. - gravity_curriculum: bool = True + Returns: + Scale tag used in cache filenames. + """ + scale_value = float(scale_value) + if scale_value <= 0.0: + raise ValueError(f"scale values must be positive, got {scale_value}") + return f"{scale_value:g}" - debug_show_axes: bool = False +def resolve_grasp_cache_file(grasp_cache_path: str, scale_value: float) -> Path: + """Resolve the grasp cache path for a single object scale. -def format_scale_tag(scale_range: Sequence[float]) -> str: - if len(scale_range) != 3: - raise ValueError(f"scale_range must have 3 values [lower, upper, num], got {scale_range}") - return f"{float(scale_range[0]):g}-{float(scale_range[1]):g}-{int(scale_range[2])}" + Args: + grasp_cache_path: Configured cache prefix or template path. + scale_value: Single object scale value for this cache file. + Returns: + Cache path for that exact scale. + """ + scale_tag = format_scale_tag(scale_value) + if "{scale}" in grasp_cache_path: + return Path(grasp_cache_path.format(scale=scale_tag)) -def resolve_grasp_cache_file(grasp_cache_path: str, scale_range: Sequence[float]) -> Path: base = Path(grasp_cache_path) if base.suffix == ".npy": - return base - return Path(f"{grasp_cache_path}_{format_scale_tag(scale_range)}.npy") + return base.with_name(f"{base.stem}_{scale_tag}{base.suffix}") + return Path(f"{grasp_cache_path}_{scale_tag}.npy") -def sample_bucketed_grasp_cache( - grasp_cache: np.ndarray, +def sample_scale_grasp_caches( + grasp_caches: Sequence[np.ndarray], scale_ids: np.ndarray, - num_scales: int, ) -> np.ndarray: + """Sample one cached grasp per reset environment from per-scale cache files. + + Args: + grasp_caches: Cache arrays ordered the same way as env scale ids. + scale_ids: Scale-bucket assignment for each reset environment. + + Returns: + Cached grasp states with shape ``(num_envs, 29)``. + """ num_envs = scale_ids.shape[0] + num_scales = len(grasp_caches) if num_scales <= 0: - raise ValueError(f"num_scales must be positive, got {num_scales}") - if grasp_cache.shape[1] < 29: - raise ValueError(f"Expected cached grasp shape (?, 29), got {grasp_cache.shape}") - if grasp_cache.shape[0] % num_scales != 0: - raise ValueError( - f"grasp_cache rows {grasp_cache.shape[0]} not divisible by num_scales={num_scales}" - ) + raise ValueError("grasp_caches must contain at least one scale bucket") - bucket = grasp_cache.shape[0] // num_scales sampled = np.zeros((num_envs, 29), dtype=np.float64) - for scale_idx in range(num_scales): + for scale_idx, grasp_cache in enumerate(grasp_caches): + if grasp_cache.ndim != 2 or grasp_cache.shape[1] < 29: + raise ValueError(f"Expected cached grasp shape (?, 29), got {grasp_cache.shape}") + if grasp_cache.shape[0] == 0: + raise ValueError(f"grasp cache for scale id {scale_idx} is empty") env_ids = np.flatnonzero(scale_ids == scale_idx) if len(env_ids) == 0: continue - sample_ids = np.random.randint(0, bucket, size=len(env_ids)) + scale_idx * bucket + sample_ids = np.random.randint(0, grasp_cache.shape[0], size=len(env_ids)) sampled[env_ids] = grasp_cache[sample_ids] return sampled @@ -240,17 +275,10 @@ def repeat_obs_history(init_frame: np.ndarray, history_len: int) -> np.ndarray: return np.asarray(history, dtype=init_frame.dtype) -def apply_random_rotation_to_positions( - positions: np.ndarray, - center: np.ndarray, - random_quat: np.ndarray, -) -> np.ndarray: - rotated = np_quat_apply(random_quat, positions - center) - return np.asarray(rotated + center, dtype=positions.dtype) - - class SharpaInhandBaseEnv(NpEnv): _cfg: SharpaInhandBaseCfg + _default_p_gain: np.ndarray + _default_d_gain: np.ndarray def __init__(self, cfg: SharpaInhandBaseCfg, backend: SimBackend, num_envs: int = 1) -> None: super().__init__(cfg, backend, num_envs) @@ -264,8 +292,16 @@ def __init__(self, cfg: SharpaInhandBaseCfg, backend: SimBackend, num_envs: int f"Model has {actuator_range.shape[0]} actuators, but Sharpa task needs {self._num_action}" ) + # Keep the raw XML actuator limits for observation normalization and any + # logic that needs the original backend contract. Target-position clipping + # uses a separate scaled limit pair to match Sharpa source behavior. self._ctrl_lower = np.asarray(actuator_range[: self._num_action, 0], dtype=self._np_dtype) self._ctrl_upper = np.asarray(actuator_range[: self._num_action, 1], dtype=self._np_dtype) + self._target_lower, self._target_upper = self._resolve_target_joint_limits( + self._ctrl_lower, + self._ctrl_upper, + cfg.control_config.dof_limits_scale, + ) self._init_qpos = self._resolve_init_qpos() self._init_qvel = np.asarray(self._backend.get_init_qvel(), dtype=np.float64) @@ -303,6 +339,7 @@ def __init__(self, cfg: SharpaInhandBaseCfg, backend: SimBackend, num_envs: int self._num_tactile = len(cfg.fingertip_body_names) self.last_contacts = np.zeros((num_envs, self._num_tactile), dtype=self._np_dtype) + self._prev_tactile_force = np.zeros((num_envs, self._num_tactile), dtype=self._np_dtype) self.object_default_pose = np.zeros((num_envs, 7), dtype=self._np_dtype) @@ -315,14 +352,42 @@ def __init__(self, cfg: SharpaInhandBaseCfg, backend: SimBackend, num_envs: int self.critic_info_buf = np.zeros((num_envs, cfg.critic_info_dim), dtype=self._np_dtype) self.scale_ids, self._num_scales, self._bucket_env = self._build_scale_ids( - num_envs, cfg.scale_range + num_envs, cfg.domain_rand.scale_list ) - self.scale_values = self._build_scale_values(cfg.scale_range) + self.scale_values = self._build_scale_values(cfg.domain_rand.scale_list) @property def action_space(self) -> gym.spaces.Box: return self._action_space # type: ignore[no-any-return] + def _resolve_target_joint_limits( + self, + raw_lower: np.ndarray, + raw_upper: np.ndarray, + scale: float, + ) -> tuple[np.ndarray, np.ndarray]: + """Build scaled target-position clipping limits from raw XML actuator bounds. + + Args: + raw_lower: Original lower control limits loaded from the backend XML. + raw_upper: Original upper control limits loaded from the backend XML. + scale: Multiplicative scale applied to both bound arrays. + + Returns: + Tuple of scaled ``(lower, upper)`` target-position limits. + """ + scale_value = float(scale) + if scale_value <= 0.0: + raise ValueError(f"dof_limits_scale must be positive, got {scale_value}") + target_lower = np.asarray(raw_lower, dtype=self._np_dtype) * scale_value + target_upper = np.asarray(raw_upper, dtype=self._np_dtype) * scale_value + if np.any(target_lower > target_upper): + raise ValueError("Scaled Sharpa target joint limits are invalid") + return target_lower.astype(self._np_dtype, copy=False), target_upper.astype( + self._np_dtype, + copy=False, + ) + def _resolve_init_qpos(self) -> np.ndarray: for key_name in ("home", "stand", "default"): try: @@ -330,53 +395,62 @@ def _resolve_init_qpos(self) -> np.ndarray: except Exception: continue - model = self._backend.model - if hasattr(model, "qpos0"): - return np.asarray(model.qpos0, dtype=np.float64) - if hasattr(model, "compute_init_dof_pos"): - return np.asarray(model.compute_init_dof_pos(), dtype=np.float64) - - raise ValueError("Could not resolve initial qpos from backend keyframes/model") + try: + return np.asarray(self._backend.get_default_qpos(), dtype=np.float64) + except NotImplementedError as exc: + raise ValueError("Could not resolve initial qpos from backend contract") from exc def _build_scale_ids( - self, num_envs: int, scale_range: Sequence[float] + self, num_envs: int, scale_list: Sequence[float] ) -> tuple[np.ndarray, int, int]: - num_scales = int(scale_range[2]) + """Build deterministic near-even environment assignments for each scale. + + Args: + num_envs: Number of vectorized environments to assign. + scale_list: Explicit object scale values used by this env instance. + + Returns: + Tuple of scale id per environment, total scale count, and the minimum + number of environments assigned to any scale. + """ + if len(scale_list) == 0: + raise ValueError("scale_list must contain at least one scale") + scale_values = np.asarray(scale_list, dtype=np.float64) + if np.any(scale_values <= 0.0): + raise ValueError(f"scale_list values must be positive, got {list(scale_list)}") + num_scales = int(scale_values.shape[0]) if num_scales <= 0: - raise ValueError(f"scale_range[2] must be >= 1, got {scale_range[2]}") - if num_envs % num_scales != 0: - raise ValueError( - f"num_envs ({num_envs}) must be divisible by scale count ({num_scales})" - ) + raise ValueError(f"scale_list must contain at least one value, got {list(scale_list)}") bucket_env = num_envs // num_scales - scale_ids = np.repeat(np.arange(num_scales, dtype=np.int32), bucket_env) + remainder = num_envs % num_scales + # Assign the remainder to the lowest scale ids to keep bucket sizes within one env. + counts = np.full((num_scales,), bucket_env, dtype=np.int32) + counts[:remainder] += 1 + scale_ids = np.repeat(np.arange(num_scales, dtype=np.int32), counts) return scale_ids, num_scales, bucket_env - def _build_scale_values(self, scale_range: Sequence[float]) -> np.ndarray: - lower = float(scale_range[0]) - upper = float(scale_range[1]) - if lower <= 0.0 or upper <= 0.0: - raise ValueError(f"scale_range bounds must be positive, got {scale_range[:2]}") - return np.asarray(np.linspace(lower, upper, self._num_scales), dtype=np.float64) + def _build_scale_values(self, scale_list: Sequence[float]) -> np.ndarray: + """Normalize configured scale values into a stable numpy array. - def _resolve_object_geom_base_size(self) -> np.ndarray | None: - if getattr(self._backend, "backend_type", None) != "mujoco": - return None + Args: + scale_list: Explicit list of object scales. - import mujoco + Returns: + Array of configured scale values in config order. + """ + scale_values = np.asarray(scale_list, dtype=np.float64) + if scale_values.ndim != 1 or scale_values.size == 0: + raise ValueError(f"scale_list must be a non-empty flat list, got {list(scale_list)}") + if np.any(scale_values <= 0.0): + raise ValueError(f"scale_list values must be positive, got {list(scale_list)}") + return scale_values - geom_id = mujoco.mj_name2id( - self._backend.model, - mujoco.mjtObj.mjOBJ_GEOM, - self._cfg.object_geom_name, - ) - if geom_id < 0: - raise ValueError(f"Geom '{self._cfg.object_geom_name}' not found in MuJoCo model") - return cast( - np.ndarray, - np.asarray(self._backend.model.geom_size[geom_id], dtype=np.float64).copy(), - ) + def _resolve_object_geom_base_size(self) -> np.ndarray | None: + try: + return cast(np.ndarray, self._backend.get_geom_size(self._cfg.object_geom_name)) + except NotImplementedError: + return None def apply_action(self, actions: np.ndarray, state: NpEnvState) -> np.ndarray: clipped_actions = np.clip(actions, -self._cfg.clip_actions, self._cfg.clip_actions) @@ -390,7 +464,9 @@ def apply_action(self, actions: np.ndarray, state: NpEnvState) -> np.ndarray: np.broadcast_to(self.default_angles, (self._num_envs, self._num_action)).copy(), ) targets = prev_targets + self._cfg.control_config.action_scale * clipped_actions - targets = np.clip(targets, self._ctrl_lower, self._ctrl_upper) + # Clip action targets by the scaled control range only. Observation + # normalization continues to use the raw XML actuator limits. + targets = np.clip(targets, self._target_lower, self._target_upper) prev_targets = np.asarray(targets, dtype=self._np_dtype) state.info["prev_targets"] = prev_targets return prev_targets @@ -427,46 +503,110 @@ def _extract_sensor_scalar(self, sensor_name: str) -> np.ndarray: flat = data.reshape(data.shape[0], -1) return np.asarray(flat[:, 0], dtype=self._np_dtype) - def _compute_tactile_observation(self) -> np.ndarray: - tactile = np.zeros((self._num_envs, self._num_tactile), dtype=self._np_dtype) - - if self._cfg.enable_tactile and self._cfg.sensor.tactile_force_sensor_names: - for i, sensor_name in enumerate( - self._cfg.sensor.tactile_force_sensor_names[: self._num_tactile] - ): - try: - tactile[:, i] = self._extract_sensor_scalar(sensor_name) - except Exception: - tactile[:, i] = 0.0 - - for disabled_id in self._cfg.disable_tactile_ids: - if 0 <= disabled_id < self._num_tactile: - tactile[:, disabled_id] = 0.0 - - latency = np.where( - np.random.rand(self._num_envs, self._num_tactile) < self._cfg.contact_latency, - 1.0, - 0.0, - ).astype(self._np_dtype) + def _read_tactile_force(self) -> np.ndarray: + """Read per-finger tactile force magnitudes in configured sensor order. - if self._cfg.binary_contact: - tactile = (tactile > self._cfg.contact_threshold).astype(self._np_dtype) - self.last_contacts = self.last_contacts * latency + tactile * (1.0 - latency) - noise_mask = ( - np.random.rand(self._num_envs, self._num_tactile) - >= self._cfg.contact_sensor_noise - ).astype(self._np_dtype) - tactile = np.where(self.last_contacts > 0.1, self.last_contacts * noise_mask, 0.0) - else: - smooth_contact = tactile * self._cfg.contact_smooth + self.last_contacts * ( - 1.0 - self._cfg.contact_smooth - ) - self.last_contacts = self.last_contacts * latency + smooth_contact * (1.0 - latency) - tactile = self.last_contacts.copy() - else: + Args: + None. + + Returns: + Array of shape ``(num_envs, num_tactile)`` ordered exactly as + ``sensor.tactile_force_sensor_names``. + """ + tactile_force = np.zeros((self._num_envs, self._num_tactile), dtype=self._np_dtype) + if not self._cfg.sensor.tactile_force_sensor_names: + return tactile_force + + for sensor_id, sensor_name in enumerate( + self._cfg.sensor.tactile_force_sensor_names[: self._num_tactile] + ): + try: + tactile_force[:, sensor_id] = self._extract_sensor_scalar(sensor_name) + except Exception: + tactile_force[:, sensor_id] = 0.0 + return tactile_force + + def _clip_tactile_force(self, tactile_force: np.ndarray) -> np.ndarray: + """Clip raw tactile-force magnitudes before they enter observation smoothing. + + Args: + tactile_force: Raw per-finger tactile magnitudes with shape + ``(num_envs, num_tactile)``. + + Returns: + Clipped tactile-force array with the same shape. Non-positive clip values + disable this clamp so the caller can opt out explicitly. + """ + clip_max = float(self._cfg.obs.tactile_force_clip_max) + tactile_force = np.asarray(tactile_force, dtype=self._np_dtype) + if clip_max <= 0.0: + return tactile_force + return np.asarray(np.clip(tactile_force, 0.0, clip_max), dtype=self._np_dtype) + + def _clear_tactile_history(self, env_ids: np.ndarray | None = None) -> None: + """Clear tactile-output and raw-force history buffers. + + Args: + env_ids: Optional environment ids to clear. When ``None``, clear all envs. + + Returns: + None. Buffers are updated in place. + """ + if env_ids is None: self.last_contacts.fill(0.0) + self._prev_tactile_force.fill(0.0) + return - return tactile + self.last_contacts[env_ids] = 0.0 + self._prev_tactile_force[env_ids] = 0.0 + + def _compute_tactile_observation(self) -> np.ndarray: + """Build tactile observations with source-equivalent smoothing and latency. + + Args: + None. + + Returns: + Tactile-force observation array with shape ``(num_envs, num_tactile)``. + """ + obs_cfg = self._cfg.obs + domain_rand = self._cfg.domain_rand + if not obs_cfg.enable_tactile: + self._clear_tactile_history() + return np.zeros((self._num_envs, self._num_tactile), dtype=self._np_dtype) + + current_force = SharpaInhandBaseEnv._clip_tactile_force(self, self._read_tactile_force()) + smooth_contact = ( + current_force * obs_cfg.contact_smooth + + self._prev_tactile_force * (1.0 - obs_cfg.contact_smooth) + ).astype(self._np_dtype) + self._prev_tactile_force[:] = current_force + + for disabled_id in self._cfg.disable_tactile_ids: + if 0 <= disabled_id < self._num_tactile: + smooth_contact[:, disabled_id] = 0.0 + + latency = np.where( + np.random.rand(self._num_envs, self._num_tactile) < domain_rand.contact_latency, + 1.0, + 0.0, + ).astype(self._np_dtype) + + if obs_cfg.binary_contact: + binary_contact = (smooth_contact > obs_cfg.contact_threshold).astype(self._np_dtype) + self.last_contacts = self.last_contacts * latency + binary_contact * (1.0 - latency) + noise_mask = ( + np.random.rand(self._num_envs, self._num_tactile) + >= domain_rand.contact_sensor_noise + ).astype(self._np_dtype) + return np.where( + self.last_contacts > 0.1, + noise_mask * self.last_contacts, + self.last_contacts, + ) + + self.last_contacts = self.last_contacts * latency + smooth_contact * (1.0 - latency) + return self.last_contacts.copy() def _compute_contact_positions(self, tactile: np.ndarray) -> np.ndarray: del tactile @@ -490,30 +630,37 @@ def _sample_pd_scales(self, lower: float, upper: float, shape: tuple[int, int]) use_small = np.random.rand(*shape) > 0.5 return np.where(use_small, small, large).astype(self._np_dtype) + def _load_default_pd_gains(self) -> tuple[np.ndarray, np.ndarray]: + """Resolve the default per-DOF PD gains used as the randomization baseline. + + Args: + None. + + Returns: + Tuple of ``(p_gain, d_gain)`` arrays with shape ``(num_action,)``. + """ + try: + p_gain, d_gain = self._backend.get_actuator_gains() + return ( + np.asarray(p_gain[: self._num_action], dtype=self._np_dtype).copy(), + np.asarray(d_gain[: self._num_action], dtype=self._np_dtype).copy(), + ) + except NotImplementedError: + return ( + np.full((self._num_action,), self._cfg.control_config.p_gain, dtype=self._np_dtype), + np.full((self._num_action,), self._cfg.control_config.d_gain, dtype=self._np_dtype), + ) + def _resolve_pd_gains(self, info: dict[str, Any]) -> tuple[np.ndarray, np.ndarray]: p_gain = info.get( "p_gain", - np.full( - (self._num_envs, self._num_action), - self._cfg.control_config.p_gain, - dtype=self._np_dtype, - ), + np.broadcast_to(self._default_p_gain, (self._num_envs, self._num_action)).copy(), ) d_gain = info.get( "d_gain", - np.full( - (self._num_envs, self._num_action), - self._cfg.control_config.d_gain, - dtype=self._np_dtype, - ), + np.broadcast_to(self._default_d_gain, (self._num_envs, self._num_action)).copy(), ) return np.asarray(p_gain, dtype=self._np_dtype), np.asarray(d_gain, dtype=self._np_dtype) def _update_proprio_history(self, obs_history: np.ndarray) -> np.ndarray: return np.asarray(obs_history[:, -self._cfg.prop_hist_len :], dtype=self._np_dtype) - - def _rotate_axis(self, axis: np.ndarray, quat: np.ndarray) -> np.ndarray: - return np.asarray(np_quat_apply(quat, axis), dtype=self._np_dtype) - - def _rotate_quat(self, quat: np.ndarray, random_quat: np.ndarray) -> np.ndarray: - return np.asarray(np_quat_mul(random_quat, quat), dtype=self._np_dtype) diff --git a/src/unilab/envs/manipulation/sharpa_inhand/grasp_gen.py b/src/unilab/envs/manipulation/sharpa_inhand/grasp_gen.py index 4fcdc98b2..0fc90af55 100644 --- a/src/unilab/envs/manipulation/sharpa_inhand/grasp_gen.py +++ b/src/unilab/envs/manipulation/sharpa_inhand/grasp_gen.py @@ -12,6 +12,7 @@ from unilab.envs.common.rotation import np_quat_error_magnitude from unilab.envs.manipulation.sharpa_inhand.base import ( SOURCE_DEFAULT_HAND_JOINT_POS_DEG, + SharpaDomainRandConfig, resolve_grasp_cache_file, ) from unilab.envs.manipulation.sharpa_inhand.rotation import ( @@ -22,28 +23,34 @@ ) +def _default_sharpa_grasp_domain_rand() -> SharpaDomainRandConfig: + """Build the nested DR defaults used by Sharpa grasp collection. + + Returns: + Domain-randomization config matching the grasp-task owner defaults. + """ + return SharpaDomainRandConfig( + randomize_pd_gains=False, + randomize_friction=False, + randomize_com=False, + randomize_mass=True, + randomize_mass_lower=0.05, + randomize_mass_upper=0.051, + force_scale=0.0, + random_force_prob_scalar=0.0, + ) + + @dataclass class SharpaInhandRotationGraspCfg(SharpaInhandRotationCfg): max_episode_seconds: float = 3.0 # 12.0 - torque_control: bool = False reset_height_lower: float = 0.61406 reset_height_upper: float = 0.62406 reset_angle_diff: float = 30.0 / 180.0 * np.pi - reset_random_quat: bool = False grasp_cache_path: str = "" - - randomize_pd_gains: bool = False - randomize_friction: bool = False - randomize_com: bool = False - randomize_mass: bool = True - randomize_mass_lower: float = 0.05 - randomize_mass_upper: float = 0.051 - - force_scale: float = 0.0 - random_force_prob_scalar: float = 0.0 - gravity_curriculum: bool = False + domain_rand: SharpaDomainRandConfig = field(default_factory=_default_sharpa_grasp_domain_rand) reward_config: RewardConfig = field( default_factory=lambda: RewardConfig( @@ -97,16 +104,25 @@ def build_reset_plan(self, env: Any, env_ids: np.ndarray) -> ResetPlan: qpos[:, env._obj_quat_slice] = object_quat qvel = np.zeros((num_reset, env.nv), dtype=np.float64) + p_gain, d_gain = self._sample_reset_pd_gains(env, num_reset, dtype=env._np_dtype) info_updates = self._build_info_updates( env, + env_ids=env_ids, hand_qpos=hand_qpos, object_pos=object_pos, object_quat=object_quat, reset_height_lower=np.full((num_reset,), env.cfg.reset_height_lower, dtype=np.float64), reset_height_upper=np.full((num_reset,), env.cfg.reset_height_upper, dtype=np.float64), rot_axis=np.broadcast_to(env._rot_axis, (num_reset, 3)).astype(np.float64), + p_gain=p_gain, + d_gain=d_gain, + friction_scale=None, + randomized_mass=None, + randomized_com_offset=None, + gravity=None, ) + env._clear_tactile_history(env_ids) return ResetPlan( env_ids=env_ids, @@ -118,7 +134,6 @@ def build_reset_plan(self, env: Any, env_ids: np.ndarray) -> ResetPlan: @registry.env("SharpaInhandRotationGrasp", sim_backend="mujoco") -@registry.env("SharpaInhandRotationGrasp", sim_backend="motrix") class SharpaInhandRotationGraspEnv(SharpaInhandRotationEnv): _cfg: SharpaInhandRotationGraspCfg @@ -128,6 +143,12 @@ def __init__( num_envs: int = 1, backend_type: str = "motrix", ) -> None: + if cfg.domain_rand.randomize_gravity or cfg.domain_rand.randomize_gravity_direction: + raise ValueError( + "SharpaInhandRotationGrasp does not support gravity randomization; " + "disable env.domain_rand.randomize_gravity and " + "env.domain_rand.randomize_gravity_direction." + ) super().__init__( cfg, num_envs=num_envs, @@ -138,7 +159,12 @@ def __init__( self._saved_grasping_states: list[list[np.ndarray]] = [ list() for _ in range(self._num_scales) ] - self._grasp_target_per_scale = max(1, int(cfg.grasp_collection_target // self._num_scales)) + if self._num_scales != 1: + raise ValueError( + "Sharpa grasp generation now collects exactly one object scale per run; " + f"got scale_list={list(self.scale_values)}" + ) + self._grasp_target_per_scale = max(1, int(cfg.grasp_collection_target)) self._grasp_cache_saved = False self._grasp_target_reached_notified = False self._last_grasp_progress_step = -1 @@ -152,6 +178,11 @@ def __init__( "Source grasp default angle count mismatch: " f"{self._grasp_default_angles.shape[0]} vs expected {self._num_action}" ) + if np.any(np.bincount(self.scale_ids, minlength=self._num_scales) == 0): + raise ValueError( + "Sharpa grasp generation requires at least one environment for the configured scale; " + f"got num_envs={num_envs}, num_scales={self._num_scales}" + ) def apply_action(self, actions: np.ndarray, state: NpEnvState) -> np.ndarray: # Grasp-cache collection should not use policy/random actions. @@ -194,7 +225,9 @@ def _maybe_print_grasp_progress(self, force: bool = False) -> None: return total = int(sum(counts)) - per_scale = ", ".join(f"s{i}:{count}" for i, count in enumerate(counts)) + per_scale = ", ".join( + f"scale={float(self.scale_values[i]):g}:{count}" for i, count in enumerate(counts) + ) print( "[SharpaInhandRotationGrasp] " f"grasp progress total={total}/{int(self._cfg.grasp_collection_target)}, " @@ -263,18 +296,12 @@ def _collect_successful_grasps(self, env_ids: np.ndarray) -> None: output_file = resolve_grasp_cache_file( self._cfg.grasp_cache_path or "cache/sharpa_grasp_linspace", - self._cfg.scale_range, + float(self.scale_values[0]), ) output_file.parent.mkdir(parents=True, exist_ok=True) - - by_scale = [] - for bucket in self._saved_grasping_states: - if bucket: - by_scale.append(np.concatenate(bucket, axis=0)[: self._grasp_target_per_scale]) - else: - by_scale.append(np.zeros((0, 29), dtype=np.float32)) - - save_data = np.concatenate(by_scale, axis=0) + save_data = np.concatenate(self._saved_grasping_states[0], axis=0)[ + : self._grasp_target_per_scale + ] np.save(output_file, save_data) self._grasp_cache_saved = True @@ -334,8 +361,7 @@ def update_state(self, state: NpEnvState) -> NpEnvState: log["grasp/cond3"] = float(np.mean(cond3.astype(np.float32))) log["grasp/valid"] = float(np.mean(grasp_valid.astype(np.float32))) per_scale_counts = self._get_per_scale_grasp_counts() - collected = float(sum(per_scale_counts)) - log["grasp/cache_size"] = collected + log["grasp/target_cache_size"] = float(self._cfg.grasp_collection_target) for scale_idx, count in enumerate(per_scale_counts): scale_value = float(self.scale_values[scale_idx]) log[f"grasp/cache_size_scale_{scale_value:g}"] = float(count) diff --git a/src/unilab/envs/manipulation/sharpa_inhand/rotation.py b/src/unilab/envs/manipulation/sharpa_inhand/rotation.py index 781e1ecf5..a46f829bb 100644 --- a/src/unilab/envs/manipulation/sharpa_inhand/rotation.py +++ b/src/unilab/envs/manipulation/sharpa_inhand/rotation.py @@ -13,19 +13,33 @@ DomainRandomizationProvider, GeomSizeOverride, InitRandomizationPlan, + IntervalRandomizationPlan, ModelVariantSpec, ResetPlan, ) from unilab.dr.dr_utils import build_common_reset_randomization, validate_common_reset_randomization +from unilab.dr.types import ( + RESET_TERM_BODY_IPOS, + RESET_TERM_BODY_MASS, + RESET_TERM_GEOM_FRICTION, + RESET_TERM_GRAVITY, + RESET_TERM_KD, + RESET_TERM_KP, + ResetRandomizationPayload, +) from unilab.dtype_config import get_global_dtype -from unilab.envs.common.rotation import np_quat_conjugate, np_quat_mul, np_quat_to_axis_angle +from unilab.envs.common.rotation import ( + np_quat_apply, + np_quat_conjugate, + np_quat_mul, + np_quat_to_axis_angle, +) from unilab.envs.manipulation.sharpa_inhand.base import ( SharpaInhandBaseCfg, SharpaInhandBaseEnv, - apply_random_rotation_to_positions, repeat_obs_history, resolve_grasp_cache_file, - sample_bucketed_grasp_cache, + sample_scale_grasp_caches, ) @@ -48,16 +62,75 @@ class RewardConfig: @registry.envcfg("SharpaInhandRotation") @dataclass class SharpaInhandRotationCfg(SharpaInhandBaseCfg): + critic_info_dim: int = 9 reward_config: RewardConfig | None = None zero_action_test_mode: bool = False - # "full": proprio+tactile+contact-pos plus critic-info channels. - # "simple": only {joint_pos, target_joint_pos, object_pos} with history. - observation_mode: str = "full" + + +def sample_random_quaternion(num_envs: int) -> np.ndarray: + """Sample uniformly distributed random quaternions in wxyz convention. + + Args: + num_envs: Number of quaternions to sample. + + Returns: + Quaternion array with shape ``(num_envs, 4)``. + """ + u1 = np.random.rand(num_envs) + u2 = np.random.rand(num_envs) * 2.0 * np.pi + u3 = np.random.rand(num_envs) * 2.0 * np.pi + + q1 = np.sqrt(1.0 - u1) * np.sin(u2) + q2 = np.sqrt(1.0 - u1) * np.cos(u2) + q3 = np.sqrt(u1) * np.sin(u3) + q4 = np.sqrt(u1) * np.cos(u3) + + return np.stack([q4, q1, q2, q3], axis=1).astype(np.float64) class SharpaInhandRotationDRProvider(DomainRandomizationProvider): def validate(self, env: Any, capabilities: DomainRandomizationCapabilities) -> None: unsupported = validate_common_reset_randomization(env, capabilities) + domain_rand = getattr(env.cfg, "domain_rand", None) + if domain_rand is not None and getattr(domain_rand, "randomize_gravity_direction", False): + if getattr(domain_rand, "randomize_gravity", False): + raise ValueError( + "Use only one Sharpa gravity randomization mode: " + "domain_rand.randomize_gravity_direction or domain_rand.randomize_gravity" + ) + if not capabilities.supports_reset_term(RESET_TERM_GRAVITY): + unsupported = unsupported | frozenset({RESET_TERM_GRAVITY}) + if ( + domain_rand is not None + and domain_rand.force_scale > 0.0 + and not capabilities.supports_interval_body_force + ): + raise NotImplementedError( + f"{env._backend.backend_type} backend does not support interval body force perturbation" + ) + if domain_rand is not None and domain_rand.randomize_pd_gains: + if not capabilities.supports_reset_term(RESET_TERM_KP): + unsupported = unsupported | frozenset({RESET_TERM_KP}) + if not capabilities.supports_reset_term(RESET_TERM_KD): + unsupported = unsupported | frozenset({RESET_TERM_KD}) + if ( + domain_rand is not None + and domain_rand.randomize_com + and not capabilities.supports_reset_term(RESET_TERM_BODY_IPOS) + ): + unsupported = unsupported | frozenset({RESET_TERM_BODY_IPOS}) + if ( + domain_rand is not None + and domain_rand.randomize_mass + and not capabilities.supports_reset_term(RESET_TERM_BODY_MASS) + ): + unsupported = unsupported | frozenset({RESET_TERM_BODY_MASS}) + if ( + domain_rand is not None + and domain_rand.randomize_friction + and not capabilities.supports_reset_term(RESET_TERM_GEOM_FRICTION) + ): + unsupported = unsupported | frozenset({RESET_TERM_GEOM_FRICTION}) if unsupported: names = ", ".join(sorted(unsupported)) raise NotImplementedError( @@ -88,91 +161,111 @@ def build_init_randomization_plan(self, env: Any) -> InitRandomizationPlan | Non model_variants=model_variants, ) - def _load_grasp_cache(self, env: Any) -> np.ndarray: - if getattr(env, "_grasp_cache", None) is not None: - return cast(np.ndarray, env._grasp_cache) + def _load_grasp_cache(self, env: Any) -> tuple[np.ndarray, ...]: + """Load one grasp cache file for each configured object scale. - cache_file = resolve_grasp_cache_file(env.cfg.grasp_cache_path, env.cfg.scale_range) - if not cache_file.exists(): - raise RuntimeError(f"No saved grasping states found at {cache_file}") + Args: + env: Sharpa rotation env instance. - env._grasp_cache = np.load(cache_file).astype(np.float64) - return cast(np.ndarray, env._grasp_cache) + Returns: + Tuple of cache arrays ordered the same as ``env.scale_values``. + """ + if getattr(env, "_grasp_cache", None) is not None: + return cast(tuple[np.ndarray, ...], env._grasp_cache) + + grasp_caches: list[np.ndarray] = [] + missing_files: list[str] = [] + for scale_value in np.asarray(env.scale_values, dtype=np.float64): + cache_file = resolve_grasp_cache_file(env.cfg.grasp_cache_path, float(scale_value)) + if not cache_file.exists(): + missing_files.append(str(cache_file)) + continue + grasp_caches.append(np.load(cache_file).astype(np.float64)) - def _sample_random_quaternion(self, num_envs: int) -> np.ndarray: - u1 = np.random.rand(num_envs) - u2 = np.random.rand(num_envs) * 2.0 * np.pi - u3 = np.random.rand(num_envs) * 2.0 * np.pi + if missing_files: + missing = ", ".join(missing_files) + raise RuntimeError(f"Missing Sharpa grasp cache file(s): {missing}") - q1 = np.sqrt(1.0 - u1) * np.sin(u2) - q2 = np.sqrt(1.0 - u1) * np.cos(u2) - q3 = np.sqrt(u1) * np.sin(u3) - q4 = np.sqrt(u1) * np.cos(u3) + env._grasp_cache = tuple(grasp_caches) + return cast(tuple[np.ndarray, ...], env._grasp_cache) - return np.stack([q4, q1, q2, q3], axis=1).astype(np.float64) + def _sample_reset_pd_gains( + self, + env: Any, + num_reset: int, + *, + dtype: np.dtype[Any], + ) -> tuple[np.ndarray, np.ndarray]: + """Sample absolute reset-time PD gains from split-around-1 scale ranges. + + Args: + env: Sharpa rotation env instance. + num_reset: Number of environments being reset. + dtype: Output dtype for the returned gain arrays. + + Returns: + Tuple of absolute ``(p_gain, d_gain)`` arrays with shape + ``(num_reset, env._num_action)``. + """ + p_gain = np.broadcast_to(env._default_p_gain, (num_reset, env._num_action)).astype( + dtype, copy=True + ) + d_gain = np.broadcast_to(env._default_d_gain, (num_reset, env._num_action)).astype( + dtype, copy=True + ) + domain_rand = env.cfg.domain_rand + if domain_rand.randomize_pd_gains: + p_scale = env._sample_pd_scales( + domain_rand.randomize_p_gain_scale_lower, + domain_rand.randomize_p_gain_scale_upper, + shape=(num_reset, env._num_action), + ) + d_scale = env._sample_pd_scales( + domain_rand.randomize_d_gain_scale_lower, + domain_rand.randomize_d_gain_scale_upper, + shape=(num_reset, env._num_action), + ) + p_gain *= p_scale + d_gain *= d_scale + return p_gain, d_gain def _build_info_updates( self, env: Any, + env_ids: np.ndarray, hand_qpos: np.ndarray, object_pos: np.ndarray, object_quat: np.ndarray, reset_height_lower: np.ndarray, reset_height_upper: np.ndarray, rot_axis: np.ndarray, + p_gain: np.ndarray, + d_gain: np.ndarray, + friction_scale: np.ndarray | None, + randomized_mass: np.ndarray | None, + randomized_com_offset: np.ndarray | None, + gravity: np.ndarray | None, ) -> dict[str, np.ndarray]: num_reset = hand_qpos.shape[0] dtype = get_global_dtype() - p_gain = np.full((num_reset, env._num_action), env.cfg.control_config.p_gain, dtype=dtype) - d_gain = np.full((num_reset, env._num_action), env.cfg.control_config.d_gain, dtype=dtype) - if env.cfg.randomize_pd_gains: - p_scale = env._sample_pd_scales( - env.cfg.randomize_p_gain_scale_lower, - env.cfg.randomize_p_gain_scale_upper, - shape=(num_reset, env._num_action), - ) - d_scale = env._sample_pd_scales( - env.cfg.randomize_d_gain_scale_lower, - env.cfg.randomize_d_gain_scale_upper, - shape=(num_reset, env._num_action), - ) - p_gain *= p_scale - d_gain *= d_scale - - critic_info = np.zeros((num_reset, env.cfg.critic_info_dim), dtype=dtype) - if env.cfg.randomize_friction: - critic_info[:, 3] = np.random.uniform( - env.cfg.randomize_friction_scale_lower, - env.cfg.randomize_friction_scale_upper, - size=(num_reset,), - ).astype(dtype) - if env.cfg.randomize_mass: - critic_info[:, 4] = np.random.uniform( - env.cfg.randomize_mass_lower, - env.cfg.randomize_mass_upper, - size=(num_reset,), - ).astype(dtype) - if env.cfg.randomize_com: - critic_info[:, 5:8] = np.random.uniform( - env.cfg.randomize_com_lower, - env.cfg.randomize_com_upper, - size=(num_reset, 3), - ).astype(dtype) - - # NOTE: source task randomizes friction/mass/com in physics directly. - # UniLab backend contract currently does not expose those object-level mutation hooks, - # so we preserve critic-info channels but keep runtime physics mutation as TODO. + critic_info = env._build_reset_critic_info( + num_reset, + env_ids, + friction_scale=friction_scale, + randomized_mass=randomized_mass, + randomized_com_offset=randomized_com_offset, + gravity=gravity, + ).astype(dtype) tactile = np.zeros((num_reset, env._num_tactile), dtype=dtype) contact_pos = np.zeros((num_reset, env._num_tactile * 3), dtype=dtype) hand_qpos_f = hand_qpos.astype(dtype) targets = hand_qpos_f.copy() object_pos_f = object_pos.astype(dtype) - init_frame = env._build_obs_frame( + init_frame = env._build_policy_frame( dof_pos=hand_qpos_f, targets=targets, - object_pos=object_pos_f, tactile=tactile, contact_pos=contact_pos, ) @@ -181,9 +274,14 @@ def _build_info_updates( object_default_pose = np.concatenate( [object_pos_f, object_quat.astype(dtype)], axis=1 ).astype(dtype) - critic_info[:, 0:3] = object_pos_f - object_default_pose[:, 0:3] + object_pos_anchor = env._build_object_pos_anchor(object_pos_f).astype(dtype) + critic_info = env._fill_critic_info( + critic_info=critic_info, + object_pos=object_pos_f, + object_pos_anchor=object_pos_anchor, + ) - return { + info_updates = { "current_actions": np.zeros((num_reset, env._num_action), dtype=dtype), "last_actions": np.zeros((num_reset, env._num_action), dtype=dtype), "prev_targets": hand_qpos_f.copy(), @@ -192,6 +290,7 @@ def _build_info_updates( "prev_object_pos": object_pos.astype(dtype).copy(), "prev_object_quat": object_quat.astype(dtype).copy(), "object_default_pose": object_default_pose, + "object_pos_anchor": object_pos_anchor, "reset_height_lower": reset_height_lower.astype(dtype), "reset_height_upper": reset_height_upper.astype(dtype), "rot_axis": rot_axis.astype(dtype), @@ -201,6 +300,7 @@ def _build_info_updates( "obs_lag_history": obs_lag_history, "proprio_hist": env._update_proprio_history(obs_lag_history), } + return info_updates def build_reset_plan(self, env: Any, env_ids: np.ndarray) -> ResetPlan: num_reset = len(env_ids) @@ -213,12 +313,13 @@ def build_reset_plan(self, env: Any, env_ids: np.ndarray) -> ResetPlan: randomization=None, ) + friction_scale = env._sample_friction_scale(num_reset) + randomized_mass = env._sample_object_mass(num_reset) + randomized_com_offset = env._sample_object_com_offset(num_reset) + gravity = env._sample_reset_gravity(num_reset) + p_gain, d_gain = self._sample_reset_pd_gains(env, num_reset, dtype=get_global_dtype()) grasp_cache = self._load_grasp_cache(env) - sampled_pose = sample_bucketed_grasp_cache( - grasp_cache, - env.scale_ids[env_ids], - env._num_scales, - ) + sampled_pose = sample_scale_grasp_caches(grasp_cache, env.scale_ids[env_ids]) hand_qpos = sampled_pose[:, : env._num_action] object_pos = sampled_pose[:, env._num_action : env._num_action + 3] @@ -226,16 +327,6 @@ def build_reset_plan(self, env: Any, env_ids: np.ndarray) -> ResetPlan: rot_axis = np.broadcast_to(env._rot_axis, (num_reset, 3)).copy().astype(np.float64) - if env.cfg.reset_random_quat: - random_quat = self._sample_random_quaternion(num_reset) - object_pos = apply_random_rotation_to_positions( - object_pos, - center=np.zeros((num_reset, 3), dtype=np.float64), - random_quat=random_quat, - ) - object_quat = env._rotate_quat(object_quat, random_quat) - rot_axis = env._rotate_axis(rot_axis, random_quat) - qpos = np.zeros((num_reset, env.nq), dtype=np.float64) qpos[:, : env._num_action] = hand_qpos qpos[:, env._obj_pos_slice] = object_pos @@ -249,20 +340,38 @@ def build_reset_plan(self, env: Any, env_ids: np.ndarray) -> ResetPlan: info_updates = self._build_info_updates( env, + env_ids=env_ids, hand_qpos=hand_qpos, object_pos=object_pos, object_quat=object_quat, reset_height_lower=reset_height_lower, reset_height_upper=reset_height_upper, rot_axis=rot_axis, + p_gain=p_gain, + d_gain=d_gain, + friction_scale=friction_scale, + randomized_mass=randomized_mass, + randomized_com_offset=randomized_com_offset, + gravity=gravity, ) + # Match the source task by clearing any cached external object force on reset. + env._random_object_force[env_ids] = 0.0 + env._clear_tactile_history(env_ids) return ResetPlan( env_ids=env_ids, qpos=qpos, qvel=qvel, info_updates=info_updates, - randomization=build_common_reset_randomization(env, num_reset), + randomization=env._build_reset_randomization( + num_reset, + p_gain=p_gain, + d_gain=d_gain, + friction_scale=friction_scale, + randomized_mass=randomized_mass, + randomized_com_offset=randomized_com_offset, + gravity=gravity, + ), ) def build_reset_observation( @@ -272,27 +381,70 @@ def build_reset_observation( info_updates: dict[str, Any], ) -> dict[str, np.ndarray]: del env_ids + tactile, contact_pos = env._policy_frame_zeros(len(info_updates["prev_targets"])) return cast( dict[str, np.ndarray], env._compute_obs_from_inputs( info_updates, dof_pos=np.asarray(info_updates["prev_targets"]), object_pos=np.asarray(info_updates["prev_object_pos"]), - tactile=np.zeros( - (len(info_updates["prev_targets"]), env._num_tactile), dtype=env._np_dtype - ), - contact_pos=np.zeros( - (len(info_updates["prev_targets"]), env._num_tactile * 3), dtype=env._np_dtype - ), + tactile=tactile, + contact_pos=contact_pos, ), ) + def build_interval_randomization_plan( + self, + env: Any, + step_counter: int, + ) -> IntervalRandomizationPlan | None: + """Build Sharpa object-force perturbations for the upcoming control step. + + Args: + env: Sharpa rotation env instance. + step_counter: Global environment step counter. + + Returns: + Interval randomization plan carrying direct object-force perturbations, + or ``None`` when object-force injection is disabled. + """ + del step_counter + domain_rand = env.cfg.domain_rand + if domain_rand.force_scale <= 0.0: + return None + + decay = float( + np.power( + domain_rand.force_decay, + env.cfg.ctrl_dt / max(domain_rand.force_decay_interval, 1.0e-8), + ) + ) + env._random_object_force *= decay + + random_mask = np.random.rand(env._num_envs) < float(domain_rand.random_force_prob_scalar) + if np.any(random_mask): + object_mass = env._resolve_current_object_mass() + env._random_object_force[random_mask] = ( + np.random.randn(int(np.sum(random_mask)), 3).astype(np.float64) + * object_mass[random_mask, None] + * float(domain_rand.force_scale) + ) + + return IntervalRandomizationPlan( + body_ids=np.asarray([env._object_body_id], dtype=np.int32), + body_force=env._random_object_force[:, None, :].copy(), + ) + @registry.env("SharpaInhandRotation", sim_backend="mujoco") -@registry.env("SharpaInhandRotation", sim_backend="motrix") class SharpaInhandRotationEnv(SharpaInhandBaseEnv): _cfg: SharpaInhandRotationCfg _reward_cfg: RewardConfig + _OBS_MODE_ALIASES: dict[str, str] = { + "separated": "separated", + "flattened": "flattened", + } + _CRITIC_BASE_DIM_WITHOUT_OPTIONALS = 8 def __init__( self, @@ -316,37 +468,33 @@ def __init__( ) super().__init__(cfg, backend, num_envs) - observation_mode = str(cfg.observation_mode).strip().lower() - if observation_mode not in ("full", "simple"): - raise ValueError( - "observation_mode must be one of {'full', 'simple'}, " - f"got {cfg.observation_mode!r}" - ) - self._observation_mode = observation_mode - - expected_full_frame_dim = ( - self._num_action + self._num_action + self._num_tactile + self._num_tactile * 3 + self._observation_mode = self._resolve_observation_mode(cfg.obs.observation_mode) + expected_critic_info_dim = self._expected_critic_info_dim() + if cfg.critic_info_dim != expected_critic_info_dim: + cfg.critic_info_dim = expected_critic_info_dim + policy_frame_dim = self._policy_frame_dim() + self.obs_buf_lag_history = np.zeros( + (num_envs, cfg.obs_history_len, policy_frame_dim), dtype=self._np_dtype ) - if self._observation_mode == "full" and cfg.frame_obs_dim != expected_full_frame_dim: - raise ValueError( - "frame_obs_dim must be " - f"{expected_full_frame_dim} for current task layout, got {cfg.frame_obs_dim}" - ) - - if cfg.torque_control: + self.proprio_hist_buf = np.zeros( + (num_envs, cfg.prop_hist_len, policy_frame_dim), dtype=self._np_dtype + ) + self.critic_info_buf = np.zeros((num_envs, expected_critic_info_dim), dtype=self._np_dtype) + self._friction_geom_ids = self._resolve_friction_geom_ids() + self._base_geom_friction = self._resolve_base_geom_friction() + self._base_gravity = self._resolve_base_gravity() + self._object_body_id = self._resolve_object_body_id() + self._base_body_mass = self._resolve_base_body_mass() + self._base_body_ipos = self._resolve_base_body_ipos() + self._default_p_gain, self._default_d_gain = self._load_default_pd_gains() + self._random_object_force = np.zeros((num_envs, 3), dtype=np.float64) + + if cfg.control_config.torque_control: raise NotImplementedError( "Sharpa torque_control=True is not implemented with the current position-actuator XML setup. " - "Set env.torque_control=false. Virtual torques are still computed explicitly for reward terms." + "Set env.control_config.torque_control=false. Virtual torques are still computed explicitly for reward terms." ) - mode = str(cfg.critic_obs_mode).strip().lower() - if mode not in ("separate", "merged"): - raise ValueError( - "critic_obs_mode must be one of {'separate', 'merged'}, " - f"got {cfg.critic_obs_mode!r}" - ) - self._critic_obs_mode = "separate" if self._observation_mode == "simple" else mode - self._reward_cfg = cfg.reward_config self._zero_action_test_mode = bool(cfg.zero_action_test_mode) self._enable_reward_log = True @@ -367,55 +515,809 @@ def apply_action(self, actions: np.ndarray, state: NpEnvState) -> np.ndarray: actions_np = np.zeros_like(actions_np, dtype=self._np_dtype) return super().apply_action(actions_np, state) - @property - def _simple_frame_obs_dim(self) -> int: - return self._num_action + self._num_action + 3 + def _scale_randomization_enabled(self) -> bool: + return self._num_scales > 1 or not np.allclose(self.scale_values, self.scale_values[0]) + + def _expected_critic_info_dim(self) -> int: + priv_info_cfg = self._cfg.priv_info + dim = self._CRITIC_BASE_DIM_WITHOUT_OPTIONALS + if priv_info_cfg.include_friction_scale: + dim += 1 + if priv_info_cfg.include_gravity_direction: + dim += 3 + return dim + + def _resolve_friction_geom_ids(self) -> dict[str, np.ndarray]: + """Resolve MuJoCo geom ids touched by Sharpa friction randomization. + + Args: + None. + + Returns: + Mapping with object, elastomer, and metal collision geom id arrays. + """ + try: + object_geom_id = self._backend.get_geom_id(self._cfg.object_geom_name) + base_body_id = self._backend.get_body_id(self._cfg.base_name) + hand_body_ids = set( + int(body_id) for body_id in self._backend.get_body_subtree_ids(base_body_id) + ) + geom_body_ids = self._backend.get_geom_body_ids() + geom_contype, geom_conaffinity = self._backend.get_geom_contact_masks() + geom_names = self._backend.get_geom_names() + except NotImplementedError: + empty = np.zeros((0,), dtype=np.int32) + return {"object": empty, "elastomer": empty, "metal": empty} + elastomer_ids: list[int] = [] + metal_ids: list[int] = [] + for geom_id, geom_name in enumerate(geom_names): + if int(geom_body_ids[geom_id]) not in hand_body_ids: + continue + if int(geom_contype[geom_id]) == 0 and int(geom_conaffinity[geom_id]) == 0: + continue + if "elastomer" in geom_name: + elastomer_ids.append(geom_id) + else: + metal_ids.append(geom_id) - def _obs_frame_dim(self) -> int: - if self._observation_mode == "simple": - return self._simple_frame_obs_dim - return int(self._cfg.frame_obs_dim) + if not elastomer_ids: + raise ValueError("No Sharpa elastomer collision geoms found for friction randomization") + if not metal_ids: + raise ValueError("No Sharpa metal collision geoms found for friction randomization") - def _build_obs_frame( + return { + "object": np.asarray([object_geom_id], dtype=np.int32), + "elastomer": np.asarray(elastomer_ids, dtype=np.int32), + "metal": np.asarray(metal_ids, dtype=np.int32), + } + + def _resolve_object_body_id(self) -> int: + """Resolve the MuJoCo body id of the manipulated object. + + Args: + None. + + Returns: + The MuJoCo object body id, or -1 when the backend is not MuJoCo. + """ + if self._object_body_ids.size == 0: + return -1 + return int(self._object_body_ids[0]) + + def _resolve_current_object_mass(self) -> np.ndarray: + """Resolve the current object mass used by force perturbation sampling. + + Args: + None. + + Returns: + Array of shape ``(num_envs,)`` with current object masses. + """ + if self.state is not None: + critic_info = np.asarray( + self.state.info.get( + "critic_info", + np.zeros((self._num_envs, self._cfg.critic_info_dim), dtype=self._np_dtype), + ), + dtype=np.float64, + ) + mass = critic_info[:, self._critic_info_layout()["mass"]].reshape(self._num_envs) + if np.all(mass > 0.0): + return mass + if self._base_body_mass is None or self._object_body_id < 0: + raise ValueError("MuJoCo object body-mass cache is unavailable") + return np.full( + (self._num_envs,), + float(self._base_body_mass[self._object_body_id]), + dtype=np.float64, + ) + + def _resolve_base_geom_friction(self) -> np.ndarray | None: + """Cache MuJoCo model friction vectors used as torsional/rolling templates. + + Args: + None. + + Returns: + Full geom friction table, or None when the backend has no geom-friction hook. + """ + try: + return np.asarray(self._backend.get_geom_friction(), dtype=np.float64).copy() + except NotImplementedError: + return None + + def _resolve_base_gravity(self) -> np.ndarray: + """Cache the default gravity vector for privileged-info fallbacks. + + Args: + None. + + Returns: + Gravity vector with shape ``(3,)``. + """ + try: + return np.asarray(self._backend.get_gravity(), dtype=np.float64).copy() + except NotImplementedError: + pass + return np.asarray([0.0, 0.0, -9.81], dtype=np.float64) + + def _resolve_base_body_mass(self) -> np.ndarray | None: + """Cache the MuJoCo body-mass table used as reset randomization baseline. + + Args: + None. + + Returns: + Full body-mass table, or None when the backend is not MuJoCo. + """ + try: + return np.asarray(self._backend.get_body_mass(), dtype=np.float64).copy() + except NotImplementedError: + return None + + def _resolve_base_body_ipos(self) -> np.ndarray | None: + """Cache the MuJoCo inertial-position table used for object COM randomization. + + Args: + None. + + Returns: + Full body inertial-position table, or None when the backend is not MuJoCo. + """ + try: + return np.asarray(self._backend.get_body_ipos(), dtype=np.float64).copy() + except NotImplementedError: + return None + + def _sample_friction_scale(self, batch_size: int) -> np.ndarray | None: + """Sample one friction multiplier per reset environment. + + Args: + batch_size: Number of reset environments. + + Returns: + Shape (batch_size, 1) multipliers, or None when disabled. + """ + domain_rand = self._cfg.domain_rand + if not domain_rand.randomize_friction: + return None + return np.random.uniform( + domain_rand.randomize_friction_scale_lower, + domain_rand.randomize_friction_scale_upper, + size=(batch_size, 1), + ).astype(np.float64) + + def _sample_object_mass(self, batch_size: int) -> np.ndarray | None: + """Sample reset-time object masses. + + Args: + batch_size: Number of reset environments. + + Returns: + Shape (batch_size, 1) absolute object masses, or None when disabled. + """ + domain_rand = self._cfg.domain_rand + if not domain_rand.randomize_mass: + return None + return np.random.uniform( + domain_rand.randomize_mass_lower, + domain_rand.randomize_mass_upper, + size=(batch_size, 1), + ).astype(np.float64) + + def _sample_object_com_offset(self, batch_size: int) -> np.ndarray | None: + """Sample reset-time object COM offsets. + + Args: + batch_size: Number of reset environments. + + Returns: + Shape (batch_size, 3) COM offsets, or None when disabled. + """ + domain_rand = self._cfg.domain_rand + if not domain_rand.randomize_com: + return None + return np.random.uniform( + domain_rand.randomize_com_lower, + domain_rand.randomize_com_upper, + size=(batch_size, 3), + ).astype(np.float64) + + def _friction_profile(self, material: str, base_sliding_friction: float) -> np.ndarray: + """Build one MuJoCo friction vector while preserving XML friction ratios. + + Args: + material: One of object, elastomer, or metal. + base_sliding_friction: Sliding-friction coefficient before random scaling. + + Returns: + Shape (3,) MuJoCo friction vector. + """ + if self._base_geom_friction is None: + raise ValueError("MuJoCo base geom friction cache is unavailable") + + geom_ids = self._friction_geom_ids[material] + template = np.asarray(self._base_geom_friction[int(geom_ids[0])], dtype=np.float64) + sliding = float(template[0]) + if sliding <= 0.0: + raise ValueError(f"{material} base sliding friction must be positive, got {sliding}") + return np.asarray(template / sliding * base_sliding_friction, dtype=np.float64) + + def _sample_reset_gravity(self, batch_size: int) -> np.ndarray | None: + """Sample the exact reset-time gravity vector applied by Sharpa. + + Args: + batch_size: Number of reset environments. + + Returns: + Gravity array with shape ``(batch_size, 3)``, or ``None`` when reset-time + gravity randomization is disabled. + """ + gravity = self._build_gravity_direction_randomization(batch_size) + if gravity is not None: + return gravity + domain_rand = self._cfg.domain_rand + if not getattr(domain_rand, "randomize_gravity", False): + return None + gravity_range = np.asarray(domain_rand.gravity_range, dtype=np.float64) + if gravity_range.shape != (2, 3): + raise ValueError( + f"domain_rand.gravity_range must have shape (2, 3), got {gravity_range.shape}" + ) + low = np.minimum(gravity_range[0], gravity_range[1]) + high = np.maximum(gravity_range[0], gravity_range[1]) + return np.random.uniform(low=low, high=high, size=(batch_size, 3)).astype(np.float64) + + def _resolved_privileged_gravity( self, - dof_pos: np.ndarray, + batch_size: int, + gravity: np.ndarray | None, + ) -> np.ndarray: + """Resolve the gravity vector written into privileged info for one reset batch. + + Args: + batch_size: Number of reset environments. + gravity: Sampled reset gravity, if reset-time randomization is active. + + Returns: + Gravity array with shape ``(batch_size, 3)``. + """ + if gravity is not None: + return np.asarray(gravity, dtype=self._np_dtype) + return np.broadcast_to(self._base_gravity, (batch_size, 3)).astype( + self._np_dtype, copy=True + ) + + def _build_friction_randomization( + self, + batch_size: int, + friction_scale: np.ndarray | None, + ) -> np.ndarray | None: + """Build the reset-time MuJoCo geom_friction randomization table. + + Args: + batch_size: Number of reset environments. + friction_scale: Shape (batch_size, 1) sampled friction multipliers. + + Returns: + Shape (batch_size, ngeom, 3) friction table, or None when disabled. + """ + if friction_scale is None: + return None + if self._base_geom_friction is None: + raise ValueError("MuJoCo base geom friction cache is unavailable") + + geom_friction = np.broadcast_to( + self._base_geom_friction, + (batch_size, *self._base_geom_friction.shape), + ).copy() + scale = np.asarray(friction_scale, dtype=np.float64).reshape(batch_size, 1, 1) + domain_rand = self._cfg.domain_rand + material_profiles = { + "object": self._friction_profile("object", domain_rand.object_base_friction), + "elastomer": self._friction_profile("elastomer", domain_rand.elastomer_base_friction), + "metal": self._friction_profile("metal", domain_rand.metal_base_friction), + } + for material, profile in material_profiles.items(): + geom_friction[:, self._friction_geom_ids[material], :] = scale * profile.reshape( + 1, 1, 3 + ) + return geom_friction + + def _build_object_mass_randomization( + self, + batch_size: int, + randomized_mass: np.ndarray | None, + ) -> np.ndarray | None: + """Build the reset-time MuJoCo body_mass table for object mass randomization. + + Args: + batch_size: Number of reset environments. + randomized_mass: Shape (batch_size, 1) sampled absolute object masses. + + Returns: + Shape (batch_size, nbody) mass table, or None when disabled. + """ + if randomized_mass is None: + return None + if self._base_body_mass is None or self._object_body_id < 0: + raise ValueError("MuJoCo base body-mass cache is unavailable") + body_mass = np.broadcast_to( + self._base_body_mass, (batch_size, self._base_body_mass.size) + ).copy() + body_mass[:, self._object_body_id] = np.asarray(randomized_mass, dtype=np.float64).reshape( + batch_size + ) + return body_mass + + def _build_object_com_randomization( + self, + batch_size: int, + randomized_com_offset: np.ndarray | None, + ) -> np.ndarray | None: + """Build the reset-time MuJoCo body_ipos table for object COM randomization. + + Args: + batch_size: Number of reset environments. + randomized_com_offset: Shape (batch_size, 3) sampled COM offsets. + + Returns: + Shape (batch_size, nbody, 3) inertial-position table, or None when disabled. + """ + if randomized_com_offset is None: + return None + if self._base_body_ipos is None or self._object_body_id < 0: + raise ValueError("MuJoCo base body-ipos cache is unavailable") + body_ipos = np.broadcast_to( + self._base_body_ipos, + (batch_size, *self._base_body_ipos.shape), + ).copy() + body_ipos[:, self._object_body_id, :] += np.asarray(randomized_com_offset, dtype=np.float64) + return body_ipos + + def _build_gravity_direction_randomization(self, batch_size: int) -> np.ndarray | None: + """Sample fixed-magnitude gravity vectors with randomized directions. + + Args: + batch_size: Number of reset environments. + + Returns: + Gravity vectors with shape ``(batch_size, 3)``, or None when disabled. + """ + domain_rand = self._cfg.domain_rand + if not getattr(domain_rand, "randomize_gravity_direction", False): + return None + magnitude = float(getattr(domain_rand, "gravity_direction_magnitude", 9.81)) + if magnitude <= 0.0: + raise ValueError(f"gravity_direction_magnitude must be positive, got {magnitude}") + gravity = np.zeros((batch_size, 3), dtype=np.float64) + gravity[:, 2] = -magnitude + random_quat = sample_random_quaternion(batch_size) + # A uniform quaternion and its inverse have the same distribution, so + # rotating gravity samples the same relative gravity directions induced + # by uniformly rotating the global hand/object/task frame. + return np.asarray(np_quat_apply(random_quat, gravity), dtype=np.float64) + + def _build_reset_randomization( + self, + batch_size: int, + *, + p_gain: np.ndarray | None, + d_gain: np.ndarray | None, + friction_scale: np.ndarray | None, + randomized_mass: np.ndarray | None, + randomized_com_offset: np.ndarray | None, + gravity: np.ndarray | None, + ) -> ResetRandomizationPayload | None: + """Build reset-randomization payloads owned by the Sharpa rotation env. + + Args: + batch_size: Number of reset environments. + friction_scale: Shape (batch_size, 1) sampled friction multipliers. + + Returns: + Reset-randomization payload, or None when no backend randomization is requested. + """ + payload = build_common_reset_randomization(self, batch_size) + if self._cfg.domain_rand.randomize_pd_gains: + if payload is None: + payload = ResetRandomizationPayload() + payload.kp = np.asarray(p_gain, dtype=np.float64) + payload.kd = np.asarray(d_gain, dtype=np.float64) + body_mass = self._build_object_mass_randomization(batch_size, randomized_mass) + body_ipos = self._build_object_com_randomization(batch_size, randomized_com_offset) + geom_friction = self._build_friction_randomization(batch_size, friction_scale) + if body_mass is not None: + if payload is None: + payload = ResetRandomizationPayload() + payload.body_mass = body_mass + if body_ipos is not None: + if payload is None: + payload = ResetRandomizationPayload() + payload.body_ipos = body_ipos + if geom_friction is not None: + if payload is None: + payload = ResetRandomizationPayload() + payload.geom_friction = geom_friction + if gravity is not None: + if payload is None: + payload = ResetRandomizationPayload() + payload.gravity = gravity + return payload + + def _critic_info_layout(self) -> dict[str, slice]: + """Describe the flat critic_info channel layout. + + Args: + None. + + Returns: + Mapping from logical field names to channel slices. + """ + offset = 0 + layout = {"object_pos_delta": slice(offset, offset + 3)} + offset += 3 + priv_info_cfg = self._cfg.priv_info + if priv_info_cfg.include_friction_scale: + layout["friction"] = slice(offset, offset + 1) + offset += 1 + layout["mass"] = slice(offset, offset + 1) + offset += 1 + layout["com"] = slice(offset, offset + 3) + offset += 3 + layout["scale"] = slice(offset, offset + 1) + offset += 1 + if priv_info_cfg.include_gravity_direction: + layout["gravity"] = slice(offset, offset + 3) + return layout + + def _assign_critic_info_field( + self, + critic_info: np.ndarray, + field_name: str, + values: np.ndarray, + ) -> None: + """Assign one critic_info field according to the declared layout. + + Args: + critic_info: Critic-info buffer to update in place. + field_name: Logical field name from the layout. + values: Batch-major values for the field. + + Returns: + None. The critic_info array is updated in place. + """ + field_slice = self._critic_info_layout().get(field_name) + if field_slice is None: + return + field_values = np.asarray(values, dtype=self._np_dtype).reshape(critic_info.shape[0], -1) + critic_info[:, field_slice] = field_values + + def _build_reset_critic_info( + self, + batch_size: int, + env_ids: np.ndarray, + *, + friction_scale: np.ndarray | None, + randomized_mass: np.ndarray | None, + randomized_com_offset: np.ndarray | None, + gravity: np.ndarray | None, + ) -> np.ndarray: + """Build reset-time critic_info for all randomized object properties. + + Args: + batch_size: Number of reset environments. + env_ids: Global environment ids for this reset batch. + + Returns: + Critic-info tensor with shape (batch_size, critic_info_dim). + """ + critic_info = np.zeros((batch_size, self._cfg.critic_info_dim), dtype=self._np_dtype) + + priv_info_cfg = self._cfg.priv_info + if priv_info_cfg.include_friction_scale: + self._assign_critic_info_field( + critic_info, + "friction", + ( + friction_scale + if friction_scale is not None + else np.ones((batch_size, 1), dtype=np.float64) + ), + ) + if randomized_mass is not None: + self._assign_critic_info_field( + critic_info, + "mass", + randomized_mass, + ) + if randomized_com_offset is not None: + self._assign_critic_info_field( + critic_info, + "com", + randomized_com_offset, + ) + if self._scale_randomization_enabled(): + self._assign_critic_info_field( + critic_info, + "scale", + self.scale_values[self.scale_ids[env_ids]].reshape(batch_size, 1), + ) + if priv_info_cfg.include_gravity_direction: + self._assign_critic_info_field( + critic_info, + "gravity", + self._resolved_privileged_gravity(batch_size, gravity), + ) + return critic_info + + def _policy_frame_dim(self) -> int: + obs_cfg = self._cfg.obs + dim = self._num_action + self._num_action + if obs_cfg.enable_tactile: + dim += self._num_tactile + if obs_cfg.enable_contact_pos: + dim += self._num_tactile * 3 + return dim + + def _policy_frame_zeros(self, batch_size: int) -> tuple[np.ndarray, np.ndarray]: + """Build zero-filled optional policy inputs for reset observations. + + Args: + batch_size: Number of environments in the batch. + + Returns: + Tuple of tactile and contact-position arrays sized for the current config. + """ + obs_cfg = self._cfg.obs + tactile_dim = self._num_tactile if obs_cfg.enable_tactile else 0 + contact_pos_dim = self._num_tactile * 3 if obs_cfg.enable_contact_pos else 0 + return ( + np.zeros((batch_size, tactile_dim), dtype=self._np_dtype), + np.zeros((batch_size, contact_pos_dim), dtype=self._np_dtype), + ) + + def _fixed_default_object_pose(self, batch_size: int) -> np.ndarray: + """Build the fixed default object pose from the backend init state. + + Args: + batch_size: Number of environments in the batch. + + Returns: + Array with shape ``(batch_size, 7)`` containing the default object + position and quaternion loaded from the backend/model init qpos. + """ + object_pos = np.broadcast_to( + np.asarray(self._init_qpos[self._obj_pos_slice], dtype=self._np_dtype), + (batch_size, 3), + ).copy() + object_quat = np.broadcast_to( + np.asarray(self._init_qpos[self._obj_quat_slice], dtype=self._np_dtype), + (batch_size, 4), + ).copy() + return np.concatenate([object_pos, object_quat], axis=1).astype(self._np_dtype) + + def _build_object_pos_anchor( + self, + object_pos: np.ndarray, + ) -> np.ndarray: + """Choose the reward/privileged-info object-position anchor. + + Args: + object_pos: Reset-time object positions with shape ``(batch_size, 3)``. + + Returns: + Array with shape ``(batch_size, 3)``. When the config flag is set, + the fixed XML/default object position is used; otherwise the sampled + reset position is returned to preserve the legacy UniLab behavior. + """ + object_pos = np.asarray(object_pos, dtype=self._np_dtype) + if self._cfg.use_default_object_pose_for_object_pos_anchor: + return self._fixed_default_object_pose(object_pos.shape[0])[:, 0:3] + return object_pos.copy() + + def _resolve_object_pos_anchor( + self, + info: dict[str, Any], + batch_size: int, + ) -> np.ndarray: + """Resolve the cached object-position anchor for one runtime batch. + + Args: + info: Mutable env-state info dictionary. + batch_size: Number of environments in the current batch. + + Returns: + Array with shape ``(batch_size, 3)`` used by the object-position + reward and privileged-info delta channels. + """ + cached_anchor = info.get("object_pos_anchor") + if cached_anchor is not None: + return np.asarray(cached_anchor, dtype=self._np_dtype) + + object_default_pose = info.get("object_default_pose") + if object_default_pose is not None: + return np.asarray(object_default_pose, dtype=self._np_dtype)[:, 0:3] + + if self._cfg.use_default_object_pose_for_object_pos_anchor: + return self._fixed_default_object_pose(batch_size)[:, 0:3] + return np.zeros((batch_size, 3), dtype=self._np_dtype) + + def _policy_frame_parts( + self, + dof_norm: np.ndarray, targets: np.ndarray, + tactile: np.ndarray, + contact_pos: np.ndarray, + ) -> list[np.ndarray]: + """Collect policy-frame components according to the configured observation layout. + + Args: + dof_norm: Normalized hand joint positions. + targets: Hand joint targets. + tactile: Tactile features. + contact_pos: Contact-position features. + + Returns: + Ordered list of arrays that should be concatenated into the policy frame. + """ + parts = [dof_norm, targets] + obs_cfg = self._cfg.obs + if obs_cfg.enable_tactile: + parts.append(np.asarray(tactile, dtype=self._np_dtype)) + if obs_cfg.enable_contact_pos: + parts.append(np.asarray(contact_pos, dtype=self._np_dtype)) + return parts + + def _fill_critic_info( + self, + critic_info: np.ndarray, object_pos: np.ndarray, + object_pos_anchor: np.ndarray, + ) -> np.ndarray: + """Populate privileged channels that are derived from runtime state. + + Args: + critic_info: Critic info buffer to update in place. + object_pos: Current object positions with shape (batch, 3). + object_pos_anchor: Object-position anchor with shape (batch, 3). + + Returns: + Updated critic info array with shape (batch, critic_info_dim). + """ + self._assign_critic_info_field( + critic_info, + "object_pos_delta", + object_pos - object_pos_anchor, + ) + return critic_info + + @classmethod + def _resolve_observation_mode(cls, observation_mode: str) -> str: + normalized_mode = str(observation_mode).strip().lower() + resolved_mode = cls._OBS_MODE_ALIASES.get(normalized_mode) + if resolved_mode is None: + supported_modes = "', '".join(sorted(cls._OBS_MODE_ALIASES)) + raise ValueError( + f"observation_mode must be one of '{supported_modes}', got {observation_mode!r}" + ) + return resolved_mode + + def _build_policy_frame( + self, + dof_pos: np.ndarray, + targets: np.ndarray, tactile: np.ndarray, contact_pos: np.ndarray, ) -> np.ndarray: dof_pos_f = np.asarray(dof_pos, dtype=self._np_dtype) targets_f = np.asarray(targets, dtype=self._np_dtype) - object_pos_f = np.asarray(object_pos, dtype=self._np_dtype) dof_norm = self._normalize_joint_pos(dof_pos_f) - if self._cfg.joint_noise_scale > 0.0: + joint_noise_scale = float(self._cfg.domain_rand.joint_noise_scale) + if joint_noise_scale > 0.0: dof_norm += ( np.random.uniform(-1.0, 1.0, size=dof_norm.shape).astype(self._np_dtype) - * self._cfg.joint_noise_scale - ) - - if self._observation_mode == "simple": - return np.asarray( - np.concatenate([dof_norm, targets_f, object_pos_f], axis=1), - dtype=self._np_dtype, + * joint_noise_scale ) - tactile_f = np.asarray(tactile, dtype=self._np_dtype) - contact_pos_f = np.asarray(contact_pos, dtype=self._np_dtype) - return np.asarray( - np.concatenate([dof_norm, targets_f, tactile_f, contact_pos_f], axis=1), + frame = np.asarray( + np.concatenate( + self._policy_frame_parts( + dof_norm=dof_norm, + targets=targets_f, + tactile=tactile, + contact_pos=contact_pos, + ), + axis=1, + ), dtype=self._np_dtype, ) + return self._clip_observation_values(frame) + + def _clip_observation_values(self, values: np.ndarray) -> np.ndarray: + """Clamp observation tensors to a configurable absolute value bound. + + Args: + values: Observation-like array to clip. + + Returns: + Clipped array with the same shape. Non-positive ``clip_obs`` disables the + clamp so callers can opt out when needed. + """ + clip_max = float(getattr(self._cfg, "clip_obs", 5.0)) + values = np.asarray(values, dtype=self._np_dtype) + if clip_max <= 0.0: + return values + return np.asarray(np.clip(values, -clip_max, clip_max), dtype=self._np_dtype) @property def obs_groups_spec(self) -> dict[str, int]: - base_obs_dim = self._cfg.obs_lag_steps * self._obs_frame_dim() - if self._observation_mode == "simple": - return {"obs": base_obs_dim} - if self._critic_obs_mode == "merged": - return {"obs": base_obs_dim + self._cfg.critic_info_dim} - return {"obs": base_obs_dim, "critic": base_obs_dim + self._cfg.critic_info_dim} + policy_obs_dim = self._cfg.obs_lag_steps * self._policy_frame_dim() + if self._observation_mode == "flattened": + return {"obs": policy_obs_dim + self._cfg.critic_info_dim} + return {"obs": policy_obs_dim, "critic": policy_obs_dim + self._cfg.critic_info_dim} + + def _build_critic_info( + self, + info: dict[str, Any], + batch_size: int, + object_pos: np.ndarray, + ) -> np.ndarray: + """Build privileged critic info for the current batch. + + Args: + info: Mutable state info dictionary carrying reset-time caches. + batch_size: Number of environments in the current batch. + object_pos: Current object positions with shape (batch, 3). + + Returns: + Privileged info array with shape (batch, critic_info_dim). + """ + critic_info = np.asarray( + info.get( + "critic_info", + np.zeros((batch_size, self._cfg.critic_info_dim), dtype=self._np_dtype), + ), + dtype=self._np_dtype, + ) + + object_pos_anchor = self._resolve_object_pos_anchor(info, batch_size) + critic_info = self._fill_critic_info( + critic_info=critic_info, + object_pos=object_pos, + object_pos_anchor=object_pos_anchor, + ) + + info["critic_info"] = critic_info + return critic_info + + def _pack_observations( + self, + policy_obs: np.ndarray, + critic_info: np.ndarray, + ) -> dict[str, np.ndarray]: + """Pack actor and privileged info into the env observation groups. + + Args: + policy_obs: Actor observation tensor with shape (batch, actor_dim). + critic_info: Privileged tensor with shape (batch, critic_info_dim). + + Returns: + Observation groups that satisfy the UniLab env contract. + """ + if self._observation_mode == "flattened": + flattened_obs = self._clip_observation_values( + np.concatenate([policy_obs, critic_info], axis=1).astype(self._np_dtype) + ) + return {"obs": flattened_obs} + + return { + "obs": self._clip_observation_values(policy_obs), + "critic": self._clip_observation_values( + np.concatenate([policy_obs, critic_info], axis=1).astype(self._np_dtype) + ), + } def _compute_obs_from_inputs( self, @@ -426,10 +1328,9 @@ def _compute_obs_from_inputs( contact_pos: np.ndarray, ) -> dict[str, np.ndarray]: targets = np.asarray(info.get("prev_targets", dof_pos), dtype=self._np_dtype) - frame = self._build_obs_frame( + frame = self._build_policy_frame( dof_pos=dof_pos, targets=targets, - object_pos=object_pos, tactile=tactile, contact_pos=contact_pos, ) @@ -450,33 +1351,8 @@ def _compute_obs_from_inputs( history[:, -self._cfg.obs_lag_steps :].reshape(batch_size, -1), dtype=self._np_dtype, ) - if self._observation_mode == "simple": - return {"obs": obs} - - critic_info = np.asarray( - info.get( - "critic_info", - np.zeros((batch_size, self._cfg.critic_info_dim), dtype=self._np_dtype), - ), - dtype=self._np_dtype, - ) - - object_default_pose = np.asarray( - info.get( - "object_default_pose", - np.zeros((batch_size, 7), dtype=self._np_dtype), - ), - dtype=self._np_dtype, - ) - if critic_info.shape[1] >= 3: - critic_info[:, 0:3] = object_pos - object_default_pose[:, 0:3] - - info["critic_info"] = critic_info - if self._critic_obs_mode == "merged": - merged_obs = np.concatenate([obs, critic_info], axis=1).astype(self._np_dtype) - return {"obs": merged_obs} - critic_obs = np.concatenate([obs, critic_info], axis=1).astype(self._np_dtype) - return {"obs": obs, "critic": critic_obs} + critic_info = self._build_critic_info(info, batch_size=batch_size, object_pos=object_pos) + return self._pack_observations(obs, critic_info) def _compute_reward( self, @@ -502,16 +1378,8 @@ def _compute_reward( torque_penalty = np.sum(np.square(torques), axis=1) work_penalty = np.square(np.sum(torques * dof_vel, axis=1)) - object_default_pose = np.asarray( - info.get( - "object_default_pose", - np.zeros((self._num_envs, 7), dtype=self._np_dtype), - ), - dtype=self._np_dtype, - ) - object_pos_reward = 1.0 / ( - np.linalg.norm(object_pos - object_default_pose[:, 0:3], axis=1) + 0.001 - ) + object_pos_anchor = self._resolve_object_pos_anchor(info, object_pos.shape[0]) + object_pos_reward = 1.0 / (np.linalg.norm(object_pos - object_pos_anchor, axis=1) + 0.001) reward_terms: dict[str, np.ndarray] = { "rotate": np.asarray(rotate_reward, dtype=self._np_dtype), @@ -539,7 +1407,7 @@ def _compute_reward( log["reward/total"] = float(np.mean(reward)) info["log"] = log - return np.asarray(reward, dtype=self._np_dtype) + return np.asarray(reward, dtype=self._np_dtype) * self._cfg.ctrl_dt def update_state(self, state: NpEnvState) -> NpEnvState: dof_pos = self.get_hand_dof_pos() diff --git a/src/unilab/ipc/shared_onpolicy_storage.py b/src/unilab/ipc/shared_onpolicy_storage.py index 34f55346f..eaea81582 100644 --- a/src/unilab/ipc/shared_onpolicy_storage.py +++ b/src/unilab/ipc/shared_onpolicy_storage.py @@ -55,7 +55,14 @@ def __init__( fields_to_allocate.pop("last_critic", None) for field, shape_fn in fields_to_allocate.items(): - shape = shape_fn(num_slots, num_envs, num_steps, obs_dim, action_dim, critic_dim) + shape = shape_fn( + num_slots, + num_envs, + num_steps, + obs_dim, + action_dim, + critic_dim, + ) nbytes = int(np.prod(shape)) * np.dtype(np.float32).itemsize if create: diff --git a/src/unilab/training/rsl_rl.py b/src/unilab/training/rsl_rl.py index 9a8f47fde..24eb585ba 100644 --- a/src/unilab/training/rsl_rl.py +++ b/src/unilab/training/rsl_rl.py @@ -134,7 +134,12 @@ def _policy_obs(self, obs: dict[str, Any]) -> torch.Tensor: return to_torch(policy_groups[0], self.device) return to_torch(np.concatenate(policy_groups, axis=1), self.device) - def _obs_to_tensordict(self, obs: dict[str, Any]) -> TensorDict: + def _obs_to_tensordict( + self, + obs: dict[str, Any], + info: dict[str, Any] | None = None, + ) -> TensorDict: + del info actor_obs = to_torch(obs["obs"], self.device) td_dict: dict[str, torch.Tensor] = { "actor": actor_obs, @@ -153,13 +158,16 @@ def _resolve_final_observation(self, state: NpEnvState) -> dict[str, Any] | None return final_observation return None + def _resolve_done(self, state: NpEnvState) -> torch.Tensor: + return to_torch(state.terminated | state.truncated, self.device).bool() + def step( self, actions: torch.Tensor | np.ndarray ) -> tuple[TensorDict, torch.Tensor, torch.Tensor, dict]: actions_np = to_numpy(actions) state = self.env.step(actions_np) rewards = to_torch(state.reward, self.device) - dones = to_torch(state.terminated | state.truncated, self.device).bool() + dones = self._resolve_done(state) self.episode_returns += rewards self.episode_lengths += 1 @@ -186,7 +194,12 @@ def step( if "log" in state.info: infos["log"] = state.info["log"] - return self._obs_to_tensordict(state.obs), rewards, dones, infos + return ( + self._obs_to_tensordict(state.obs, getattr(state, "info", None)), + rewards, + dones, + infos, + ) def reset(self) -> tuple[TensorDict, dict[str, Any]]: if self.env.state is None: @@ -196,11 +209,11 @@ def reset(self) -> tuple[TensorDict, dict[str, Any]]: obs_out, info = self.env.reset(env_indices) self.episode_returns[:] = 0 self.episode_lengths[:] = 0 - return self._obs_to_tensordict(obs_out), info + return self._obs_to_tensordict(obs_out, info), info def get_observations(self) -> TensorDict: assert self.env.state is not None - return self._obs_to_tensordict(self.env.state.obs) + return self._obs_to_tensordict(self.env.state.obs, self.env.state.info) def get_privileged_observations(self) -> torch.Tensor: assert self.env.state is not None diff --git a/src/unilab/utils/support_matrix.py b/src/unilab/utils/support_matrix.py index 36764cfa1..d905f15b4 100644 --- a/src/unilab/utils/support_matrix.py +++ b/src/unilab/utils/support_matrix.py @@ -191,6 +191,8 @@ def _configured_entries(root: Path, spec: EntrypointSpec) -> dict[str, dict[str, for task_path in sorted(task_root.glob(spec.task_glob)): task_slug = task_path.parent.name backend = task_path.stem + if backend not in BACKENDS: + continue entries.setdefault(task_slug, {})[backend] = _load_task_name(task_path) return entries diff --git a/tests/algos/test_hora_contract.py b/tests/algos/test_hora_contract.py new file mode 100644 index 000000000..5615912dc --- /dev/null +++ b/tests/algos/test_hora_contract.py @@ -0,0 +1,44 @@ +from __future__ import annotations + +import ast +import inspect +import textwrap + +import pytest +import torch + + +def test_hora_rsl_wrapper_uses_explicit_np_env_state_contract() -> None: + """HORA wrapper must not probe required NpEnvState fields dynamically.""" + from unilab.algos.torch.hora.rsl_rl import HoraRslRlVecEnvWrapper + + source = textwrap.dedent(inspect.getsource(HoraRslRlVecEnvWrapper.step)) + tree = ast.parse(source) + forbidden_calls: list[str] = [] + for node in ast.walk(tree): + if not isinstance(node, ast.Call) or not isinstance(node.func, ast.Name): + continue + if node.func.id not in {"getattr", "hasattr"}: + continue + if not node.args or not isinstance(node.args[0], ast.Name): + continue + if node.args[0].id == "state": + forbidden_calls.append(node.func.id) + + assert forbidden_calls == [] + + +def test_hora_appo_learner_derives_priv_info_from_critic_contract() -> None: + from unilab.algos.torch.hora.appo_learner import _derive_priv_info_from_critic + + actor_obs = torch.zeros((2, 3, 4), dtype=torch.float32) + priv_info = torch.arange(12, dtype=torch.float32).reshape(2, 3, 2) + critic_obs = torch.cat([actor_obs, priv_info], dim=-1) + + torch.testing.assert_close( + _derive_priv_info_from_critic(actor_obs, critic_obs, context="test"), + priv_info, + ) + + with pytest.raises(ValueError, match="privileged tail"): + _derive_priv_info_from_critic(actor_obs, actor_obs, context="test") diff --git a/tests/algos/test_hora_imports.py b/tests/algos/test_hora_imports.py new file mode 100644 index 000000000..f5cfcadee --- /dev/null +++ b/tests/algos/test_hora_imports.py @@ -0,0 +1,51 @@ +from __future__ import annotations + +import importlib +import sys + +import pytest +from rsl_rl.utils import resolve_callable + + +def test_hora_package_import_keeps_appo_lazy() -> None: + sys.modules.pop("unilab.algos.torch.hora", None) + sys.modules.pop("unilab.algos.torch.hora.appo", None) + + importlib.import_module("unilab.algos.torch.hora") + + assert "unilab.algos.torch.hora.appo" not in sys.modules + + +def test_resolve_callable_loads_hora_ppo_from_package_export() -> None: + resolved = resolve_callable("unilab.algos.torch.hora:HoraPPO") + + from unilab.algos.torch.hora.ppo import HoraPPO + + assert resolved is HoraPPO + + +def test_rsl_rl_runtime_resolver_loads_hora_wrapper_from_owner_marker() -> None: + from unilab.algos.torch.hora.rsl_rl import HoraRslRlVecEnvWrapper + from unilab.algos.torch.rsl_rl_runtime import resolve_rsl_rl_ppo_runtime + from unilab.training.rsl_rl import RslRlVecEnvWrapper + + runtime = resolve_rsl_rl_ppo_runtime( + { + "runtime_impl": "hora_ppo", + "runtime_resolver": "unilab.algos.torch.hora.rsl_rl:resolve_hora_ppo_runtime", + }, + default_wrapper_cls=RslRlVecEnvWrapper, + ) + + assert runtime.wrapper_cls is HoraRslRlVecEnvWrapper + + +def test_rsl_rl_runtime_resolver_rejects_unresolved_custom_runtime() -> None: + from unilab.algos.torch.rsl_rl_runtime import resolve_rsl_rl_ppo_runtime + from unilab.training.rsl_rl import RslRlVecEnvWrapper + + with pytest.raises(ValueError, match="runtime_impl='hora_ppo'.*runtime_resolver"): + resolve_rsl_rl_ppo_runtime( + {"runtime_impl": "hora_ppo"}, + default_wrapper_cls=RslRlVecEnvWrapper, + ) diff --git a/tests/base/test_sim_backend_smoke.py b/tests/base/test_sim_backend_smoke.py index eff64f54a..38128cb88 100644 --- a/tests/base/test_sim_backend_smoke.py +++ b/tests/base/test_sim_backend_smoke.py @@ -29,6 +29,7 @@ def _xml(robot: str, scene: str = "scene_flat.xml") -> str: _G1 = dict(model_file=_xml("g1"), base_name="pelvis") _ALLEGRO = dict(model_file=_xml("allegro_hand", "scene.xml"), base_name="palm") +_SHARPA = dict(model_file=_xml("sharpa_wave", "scene.xml"), base_name="right_hand_C_MC") NUM_ENVS = 2 SIM_DT = 0.005 @@ -221,6 +222,77 @@ def test_mujoco_model_properties_smoke(): assert expected_nv >= bkd.num_dof_vel +def test_mujoco_metadata_getters_return_stable_copies(): + mujoco = _mujoco_module() + + from unilab.base.backend.mujoco_backend import MuJoCoBackend + + bkd = MuJoCoBackend(_SHARPA["model_file"], NUM_ENVS, SIM_DT, base_name=_SHARPA["base_name"]) + model = bkd.model + object_geom_id = int(mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_GEOM, "object")) + base_body_id = int(mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, _SHARPA["base_name"])) + + assert bkd.get_geom_id("object") == object_geom_id + assert bkd.get_body_id(_SHARPA["base_name"]) == base_body_id + with pytest.raises(ValueError, match="Geom 'missing'"): + bkd.get_geom_id("missing") + with pytest.raises(ValueError, match="Body 'missing'"): + bkd.get_body_id("missing") + + default_qpos = bkd.get_default_qpos() + _shape(default_qpos, model.nq) + np.testing.assert_allclose(default_qpos, model.qpos0) + default_qpos[0] += 1.0 + assert not np.isclose(default_qpos[0], model.qpos0[0]) + + geom_size = bkd.get_geom_size("object") + _shape(geom_size, 3) + np.testing.assert_allclose(geom_size, model.geom_size[object_geom_id]) + geom_size[0] += 1.0 + assert not np.isclose(geom_size[0], model.geom_size[object_geom_id, 0]) + + geom_body_ids = bkd.get_geom_body_ids() + _shape(geom_body_ids, model.ngeom) + np.testing.assert_array_equal(geom_body_ids, model.geom_bodyid) + geom_body_ids[object_geom_id] = -1 + assert int(model.geom_bodyid[object_geom_id]) != -1 + + geom_contype, geom_conaffinity = bkd.get_geom_contact_masks() + _shape(geom_contype, model.ngeom) + _shape(geom_conaffinity, model.ngeom) + np.testing.assert_array_equal(geom_contype, model.geom_contype) + np.testing.assert_array_equal(geom_conaffinity, model.geom_conaffinity) + + geom_names = bkd.get_geom_names() + assert len(geom_names) == model.ngeom + assert geom_names[object_geom_id] == "object" + assert base_body_id in set(int(body_id) for body_id in bkd.get_body_subtree_ids(base_body_id)) + + geom_friction = bkd.get_geom_friction() + _shape(geom_friction, model.ngeom, 3) + np.testing.assert_allclose(geom_friction, model.geom_friction) + geom_friction[object_geom_id, 0] += 1.0 + assert not np.isclose(geom_friction[object_geom_id, 0], model.geom_friction[object_geom_id, 0]) + + gravity = bkd.get_gravity() + _shape(gravity, 3) + np.testing.assert_allclose(gravity, model.opt.gravity) + gravity[2] += 1.0 + assert not np.isclose(gravity[2], model.opt.gravity[2]) + + body_mass = bkd.get_body_mass() + _shape(body_mass, model.nbody) + np.testing.assert_allclose(body_mass, model.body_mass) + body_mass[base_body_id] += 1.0 + assert not np.isclose(body_mass[base_body_id], model.body_mass[base_body_id]) + + body_ipos = bkd.get_body_ipos() + _shape(body_ipos, model.nbody, 3) + np.testing.assert_allclose(body_ipos, model.body_ipos) + body_ipos[base_body_id, 0] += 1.0 + assert not np.isclose(body_ipos[base_body_id, 0], model.body_ipos[base_body_id, 0]) + + def test_motrix_model_properties_smoke(): pytest.importorskip("motrixsim") @@ -231,4 +303,5 @@ def test_motrix_model_properties_smoke(): assert bkd.num_dof_vel > 0 ctrl_range = bkd.get_actuator_ctrl_range() _shape(ctrl_range, bkd.num_actuators, 2) + assert bkd.get_default_qpos().ndim == 1 assert bkd.get_joint_range() is None diff --git a/tests/benchmark/test_sharpa_init_dr_benchmark.py b/tests/benchmark/test_sharpa_init_dr_benchmark.py new file mode 100644 index 000000000..142aca664 --- /dev/null +++ b/tests/benchmark/test_sharpa_init_dr_benchmark.py @@ -0,0 +1,19 @@ +from __future__ import annotations + +import numpy as np +from benchmark import benchmark_sharpa_init_dr_construct as sharpa_benchmark + + +def test_sharpa_init_dr_benchmark_uses_owner_scale_list_key() -> None: + cfg = sharpa_benchmark._compose_cfg( + "sharpa_inhand/mujoco", + lower=0.5, + upper=0.8, + variant_count=4, + ) + + np.testing.assert_allclose( + np.asarray(cfg.env.domain_rand.scale_list, dtype=np.float64), + np.linspace(0.5, 0.8, 4, dtype=np.float64), + ) + assert "scale_list" not in cfg.env diff --git a/tests/config/test_config_system.py b/tests/config/test_config_system.py index a564a46ff..1fa18a0e7 100644 --- a/tests/config/test_config_system.py +++ b/tests/config/test_config_system.py @@ -18,6 +18,14 @@ CONF_DIR = Path(__file__).parent.parent.parent / "conf" _PPO_MLX_TASKS = {"go1_joystick_flat", "go2_joystick_flat", "g1_walk_flat"} +_BACKENDS = ("mujoco", "motrix") + + +def _expected_backend_from_variant(name: str) -> str | None: + for backend in _BACKENDS: + if name == backend or name.startswith(f"{backend}_"): + return backend + return None def _compose(algo_dir: str, config_name: str = "config", overrides: list[str] | None = None): @@ -68,12 +76,15 @@ def _supported_task_cases() -> list[tuple[str, str, str, str, str, list[str]]]: root = CONF_DIR / algo_dir / "task" for task_dir in sorted(path for path in root.iterdir() if path.is_dir()): for backend_file in sorted(task_dir.glob("*.yaml")): + expected_backend = _expected_backend_from_variant(backend_file.stem) + if expected_backend is None: + continue cases.append( ( algo_dir, "config", task_dir.name, - backend_file.stem, + expected_backend, str(backend_file.relative_to(CONF_DIR)), [f"task={task_dir.name}/{backend_file.stem}"], ) @@ -84,7 +95,7 @@ def _supported_task_cases() -> list[tuple[str, str, str, str, str, list[str]]]: algo_dir, "config_mlx", task_dir.name, - backend_file.stem, + expected_backend, str(backend_file.relative_to(CONF_DIR)), [f"task={task_dir.name}/{backend_file.stem}"], ) @@ -94,12 +105,15 @@ def _supported_task_cases() -> list[tuple[str, str, str, str, str, list[str]]]: for algo_root in sorted(path for path in offpolicy_root.iterdir() if path.is_dir()): for task_dir in sorted(path for path in algo_root.iterdir() if path.is_dir()): for backend_file in sorted(task_dir.glob("*.yaml")): + expected_backend = _expected_backend_from_variant(backend_file.stem) + if expected_backend is None: + continue cases.append( ( "offpolicy", "config", task_dir.name, - backend_file.stem, + expected_backend, str(backend_file.relative_to(CONF_DIR)), [ f"algo={algo_root.name}", @@ -151,6 +165,8 @@ def test_task_files_keep_full_identity_without_hidden_backend_marker(): assert "_selected_sim_backend" not in cfg_dict_raw, ( f"task has hidden backend marker: {path}" ) + if path.stem not in _BACKENDS: + continue training_raw = cfg_dict_raw.get("training", {}) assert isinstance(training_raw, dict) assert "task_name" in training_raw, f"task missing task_name: {path}" diff --git a/tests/envs/test_env_configs.py b/tests/envs/test_env_configs.py index 6f3d5a166..49a06b28f 100644 --- a/tests/envs/test_env_configs.py +++ b/tests/envs/test_env_configs.py @@ -12,6 +12,7 @@ import sys import textwrap from pathlib import Path +from types import SimpleNamespace from typing import Any, cast import numpy as np @@ -50,7 +51,7 @@ def blocked_import(name, globals=None, locals=None, fromlist=(), level=0): from unilab.base import registry from unilab.base.backend import create_backend - from unilab.envs.manipulation.inhand_rot_allegro.rotation import AllegroRotationCfg + from unilab.envs.manipulation.allegro_inhand.rotation import AllegroRotationCfg from unilab.envs.motion_tracking.g1.tracking import G1MotionTrackingCfg from unilab.base.registry import ensure_registries @@ -318,7 +319,7 @@ def test_g1_walk_flat_assets_define_contact_sensors_for_gait_rewards(): def test_allegro_rotation_obs_groups_spec_dims(): """Allegro rotation obs_groups_spec should expose single actor obs group.""" - from unilab.envs.manipulation.inhand_rot_allegro.rotation import AllegroRotationPPO + from unilab.envs.manipulation.allegro_inhand.rotation import AllegroRotationPPO env = cast(Any, object.__new__(AllegroRotationPPO)) spec = env.obs_groups_spec @@ -328,7 +329,7 @@ def test_allegro_rotation_obs_groups_spec_dims(): def test_allegro_grasp_obs_groups_spec_dims(): """Allegro grasp task inherits the same obs group layout as rotation.""" - from unilab.envs.manipulation.inhand_rot_allegro.grasp_gen import AllegroRotationGrasp + from unilab.envs.manipulation.allegro_inhand.grasp_gen import AllegroRotationGrasp env = cast(Any, object.__new__(AllegroRotationGrasp)) spec = env.obs_groups_spec @@ -470,10 +471,34 @@ def fake_base_init(self, cfg, backend, num_envs): self._np_dtype = np.float64 self._num_action = 22 self._num_tactile = 5 - self._num_scales = int(cfg.scale_range[2]) + self._num_scales = len(cfg.domain_rand.scale_list) self.scale_ids = np.zeros((num_envs,), dtype=np.int32) + self._object_body_ids = np.zeros((0,), dtype=np.int32) - monkeypatch.setattr(sharpa_rotation_module, "create_backend", lambda *args, **kwargs: object()) + def unsupported_backend_metadata(*args, **kwargs): + raise NotImplementedError("fake backend does not expose Sharpa MuJoCo metadata") + + monkeypatch.setattr( + sharpa_rotation_module, + "create_backend", + lambda *args, **kwargs: SimpleNamespace( + backend_type="motrix", + get_actuator_gains=lambda: ( + np.ones(22, dtype=np.float64), + np.ones(22, dtype=np.float64), + ), + get_geom_id=unsupported_backend_metadata, + get_body_id=unsupported_backend_metadata, + get_body_subtree_ids=unsupported_backend_metadata, + get_geom_body_ids=unsupported_backend_metadata, + get_geom_contact_masks=unsupported_backend_metadata, + get_geom_names=unsupported_backend_metadata, + get_geom_friction=unsupported_backend_metadata, + get_gravity=unsupported_backend_metadata, + get_body_mass=unsupported_backend_metadata, + get_body_ipos=unsupported_backend_metadata, + ), + ) monkeypatch.setattr(SharpaInhandBaseEnv, "__init__", fake_base_init) monkeypatch.setattr( SharpaInhandRotationGraspEnv, @@ -482,6 +507,27 @@ def fake_base_init(self, cfg, backend, num_envs): ) cfg = SharpaInhandRotationGraspCfg() + assert cfg.domain_rand.randomize_pd_gains is False + assert cfg.domain_rand.randomize_friction is False + assert cfg.domain_rand.randomize_com is False + assert cfg.domain_rand.randomize_mass is True + assert cfg.domain_rand.randomize_mass_lower == pytest.approx(0.05) + assert cfg.domain_rand.randomize_mass_upper == pytest.approx(0.051) + assert cfg.domain_rand.force_scale == pytest.approx(0.0) + assert cfg.domain_rand.random_force_prob_scalar == pytest.approx(0.0) + assert cfg.domain_rand.joint_noise_scale == pytest.approx(0.02) + assert cfg.domain_rand.contact_latency == pytest.approx(0.005) + assert cfg.domain_rand.contact_sensor_noise == pytest.approx(0.01) + assert cfg.control_config.torque_control is False + assert cfg.control_config.dof_limits_scale == pytest.approx(0.9) + assert cfg.obs.enable_tactile is True + assert cfg.obs.binary_contact is False + assert cfg.obs.enable_contact_pos is False + assert cfg.obs.contact_smooth == pytest.approx(0.5) + assert cfg.obs.contact_threshold == pytest.approx(0.05) + assert cfg.obs.tactile_force_clip_max == pytest.approx(4.0) + assert cfg.priv_info.include_friction_scale is True + assert cfg.priv_info.include_gravity_direction is False env = cast(Any, SharpaInhandRotationGraspEnv)(cfg, num_envs=4, backend_type="mujoco") assert calls == ["SharpaInhandGraspDRProvider"] diff --git a/tests/envs/test_sharpa.py b/tests/envs/test_sharpa.py new file mode 100644 index 000000000..cf0d2a049 --- /dev/null +++ b/tests/envs/test_sharpa.py @@ -0,0 +1,663 @@ +from __future__ import annotations + +from pathlib import Path +from tempfile import TemporaryDirectory +from types import SimpleNamespace +from typing import Any + +import numpy as np +import pytest +from hydra import compose, initialize_config_dir +from hydra.core.global_hydra import GlobalHydra +from omegaconf import OmegaConf + +from unilab.base.registry import ensure_registries +from unilab.envs.manipulation.sharpa_inhand.base import SharpaInhandBaseEnv +from unilab.envs.manipulation.sharpa_inhand.rotation import SharpaInhandRotationDRProvider + +_CONF_DIR = Path(__file__).resolve().parents[2] / "conf" +_SRC_DIR = Path(__file__).resolve().parents[2] / "src" + + +def test_sharpa_env_uses_backend_contract_for_mujoco_metadata() -> None: + """Sharpa env code should not read MuJoCo model internals directly.""" + source = "\n".join( + (_SRC_DIR / "unilab" / "envs" / "manipulation" / "sharpa_inhand" / path).read_text( + encoding="utf-8" + ) + for path in ("base.py", "rotation.py") + ) + + assert "import mujoco" not in source + assert "self._backend.model" not in source + assert "_backend.model" not in source + + +def _require_mujoco_runtime() -> None: + """Require the MuJoCo batch runtime used by Sharpa reset randomization tests. + + Args: + None. + + Returns: + None. The helper skips the test when MuJoCo batch runtime is unavailable. + """ + pytest.importorskip("mujoco", reason="mujoco not installed") + try: + from mujoco.batch_env import BatchEnvPool as _ # noqa: F401 + except Exception: + pytest.skip( + "mujoco.batch_env not available (platform/libstdc++ issue)", + allow_module_level=False, + ) + + +def _compose_sharpa_mujoco_owner_cfg(num_envs: int) -> tuple[Any, dict[str, Any]]: + """Compose the Sharpa MuJoCo owner config used by the real training path. + + Args: + num_envs: Number of vectorized environments for the test env. + + Returns: + Tuple of the composed Hydra config and the env_cfg_override dict for registry.make(). + """ + GlobalHydra.instance().clear() + with initialize_config_dir(config_dir=str(_CONF_DIR / "ppo"), version_base="1.3"): + cfg = compose( + "config", + overrides=[ + "task=sharpa_inhand/mujoco", + f"algo.num_envs={num_envs}", + ], + ) + + env_cfg_override = OmegaConf.to_container(cfg.env, resolve=True) + assert isinstance(env_cfg_override, dict) + env_cfg_override["reward_config"] = OmegaConf.to_container(cfg.reward, resolve=True) + return cfg, env_cfg_override + + +def _build_fake_tactile_env( + sensor_data: dict[str, np.ndarray], + *, + enable_tactile: bool = True, + binary_contact: bool = False, + disable_tactile_ids: list[int] | None = None, + contact_smooth: float = 0.5, + contact_threshold: float = 0.05, + contact_latency: float = 0.0, + contact_sensor_noise: float = 0.01, + last_contacts: np.ndarray | None = None, + prev_tactile_force: np.ndarray | None = None, +) -> Any: + """Build a minimal fake Sharpa env for tactile-observation unit tests. + + Args: + sensor_data: Mapping from tactile sensor name to backend sensor array. + enable_tactile: Whether tactile observation is enabled. + binary_contact: Whether binary-contact mode is enabled. + disable_tactile_ids: Optional tactile ids to zero out. + contact_smooth: Smoothing weight applied to the latest raw force. + contact_threshold: Threshold used by binary-contact mode. + contact_latency: Bernoulli probability of keeping the previous tactile output. + contact_sensor_noise: Binary-contact sensor dropout probability. + last_contacts: Optional previous tactile output buffer. + prev_tactile_force: Optional previous raw tactile-force buffer. + + Returns: + Simple fake env exposing the fields used by ``_compute_tactile_observation``. + """ + tactile_names = [ + "contact_right_thumb_elastomer_force", + "contact_right_index_elastomer_force", + "contact_right_middle_elastomer_force", + "contact_right_ring_elastomer_force", + "contact_right_pinky_elastomer_force", + ] + num_envs = next(iter(sensor_data.values())).shape[0] + env = SimpleNamespace( + _num_envs=num_envs, + _num_tactile=len(tactile_names), + _np_dtype=np.float64, + _backend=SimpleNamespace(get_sensor_data=lambda name: sensor_data[name]), + _cfg=SimpleNamespace( + obs=SimpleNamespace( + enable_tactile=enable_tactile, + binary_contact=binary_contact, + contact_smooth=contact_smooth, + contact_threshold=contact_threshold, + tactile_force_clip_max=5.0, + ), + domain_rand=SimpleNamespace( + contact_latency=contact_latency, + contact_sensor_noise=contact_sensor_noise, + ), + disable_tactile_ids=list(disable_tactile_ids or []), + sensor=SimpleNamespace(tactile_force_sensor_names=tactile_names), + ), + last_contacts=np.zeros((num_envs, len(tactile_names)), dtype=np.float64) + if last_contacts is None + else np.asarray(last_contacts, dtype=np.float64).copy(), + _prev_tactile_force=np.zeros((num_envs, len(tactile_names)), dtype=np.float64) + if prev_tactile_force is None + else np.asarray(prev_tactile_force, dtype=np.float64).copy(), + ) + env._extract_sensor_scalar = lambda sensor_name: SharpaInhandBaseEnv._extract_sensor_scalar( + env, sensor_name + ) + env._read_tactile_force = lambda: SharpaInhandBaseEnv._read_tactile_force(env) + env._clear_tactile_history = lambda env_ids=None: SharpaInhandBaseEnv._clear_tactile_history( + env, env_ids + ) + return env + + +def test_sharpa_provider_builds_mujoco_init_geom_scale_plan() -> None: + env = SimpleNamespace( + _backend=SimpleNamespace(backend_type="mujoco"), + _object_geom_base_size=np.array([0.02, 0.016, 0.0], dtype=np.float64), + scale_values=np.array([0.5, 0.8], dtype=np.float64), + scale_ids=np.array([0, 0, 1, 1], dtype=np.int32), + cfg=SimpleNamespace(object_geom_name="object"), + ) + + plan = SharpaInhandRotationDRProvider().build_init_randomization_plan(env) + + assert plan is not None + np.testing.assert_array_equal( + plan.model_assignments, + np.array([0, 0, 1, 1], dtype=np.int32), + ) + assert len(plan.model_variants) == 2 + assert plan.model_variants[0].geom_size_overrides[0].geom_name == "object" + np.testing.assert_allclose( + plan.model_variants[0].geom_size_overrides[0].size, + [0.01, 0.008, 0.0], + ) + np.testing.assert_allclose( + plan.model_variants[1].geom_size_overrides[0].size, + [0.016, 0.0128, 0.0], + ) + + +def test_sharpa_provider_skips_non_mujoco_init_geom_scale_plan() -> None: + env = SimpleNamespace( + _backend=SimpleNamespace(backend_type="motrix"), + _object_geom_base_size=np.array([0.02, 0.016, 0.0], dtype=np.float64), + scale_values=np.array([0.5], dtype=np.float64), + scale_ids=np.array([0], dtype=np.int32), + cfg=SimpleNamespace(object_geom_name="object"), + ) + + plan = SharpaInhandRotationDRProvider().build_init_randomization_plan(env) + + assert plan is None + + +def test_sharpa_tactile_force_matches_reference_smoothing_and_order() -> None: + """Verify tactile force uses reference sensor order and raw-force smoothing. + + Args: + None. + + Returns: + None. The assertions validate thumb→pinky ordering and 2-step smoothing. + """ + sensor_data = { + "contact_right_thumb_elastomer_force": np.array([[1.0, 0.0, 0.0]], dtype=np.float64), + "contact_right_index_elastomer_force": np.array([[2.0, 0.0, 0.0]], dtype=np.float64), + "contact_right_middle_elastomer_force": np.array([[3.0, 0.0, 0.0]], dtype=np.float64), + "contact_right_ring_elastomer_force": np.array([[4.0, 0.0, 0.0]], dtype=np.float64), + "contact_right_pinky_elastomer_force": np.array([[5.0, 0.0, 0.0]], dtype=np.float64), + } + env = _build_fake_tactile_env( + sensor_data, + contact_smooth=0.25, + prev_tactile_force=np.array([[10.0, 20.0, 30.0, 40.0, 50.0]], dtype=np.float64), + ) + + tactile = SharpaInhandBaseEnv._compute_tactile_observation(env) + + expected = np.array([[7.75, 15.5, 23.25, 31.0, 38.75]], dtype=np.float64) + np.testing.assert_allclose(tactile, expected) + np.testing.assert_allclose(env.last_contacts, expected) + np.testing.assert_allclose( + env._prev_tactile_force, + np.array([[1.0, 2.0, 3.0, 4.0, 5.0]], dtype=np.float64), + ) + + +def test_sharpa_tactile_binary_mode_matches_reference_latency_noise_and_disable( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Verify binary tactile mode matches reference latency/noise semantics. + + Args: + monkeypatch: Pytest helper used to make the random masks deterministic. + + Returns: + None. The assertions validate binary thresholding, latency, and dropout. + """ + sensor_data = { + "contact_right_thumb_elastomer_force": np.array([[0.2, 0.0, 0.0]], dtype=np.float64), + "contact_right_index_elastomer_force": np.array([[0.01, 0.0, 0.0]], dtype=np.float64), + "contact_right_middle_elastomer_force": np.array([[0.5, 0.0, 0.0]], dtype=np.float64), + "contact_right_ring_elastomer_force": np.array([[0.3, 0.0, 0.0]], dtype=np.float64), + "contact_right_pinky_elastomer_force": np.array([[0.9, 0.0, 0.0]], dtype=np.float64), + } + env = _build_fake_tactile_env( + sensor_data, + binary_contact=True, + disable_tactile_ids=[4], + contact_smooth=1.0, + contact_threshold=0.05, + contact_latency=0.5, + contact_sensor_noise=0.01, + last_contacts=np.array([[0.4, 0.7, 0.0, 1.0, 0.2]], dtype=np.float64), + ) + + sampled_masks = iter( + [ + np.array([[0.9, 0.1, 0.9, 0.9, 0.1]], dtype=np.float64), + np.array([[0.2, 0.2, 0.005, 0.9, 0.2]], dtype=np.float64), + ] + ) + + monkeypatch.setattr(np.random, "rand", lambda *shape: next(sampled_masks).copy()) + + tactile = SharpaInhandBaseEnv._compute_tactile_observation(env) + + expected_last = np.array([[1.0, 0.7, 1.0, 1.0, 0.2]], dtype=np.float64) + expected_tactile = np.array([[1.0, 0.7, 0.0, 1.0, 0.2]], dtype=np.float64) + np.testing.assert_allclose(env.last_contacts, expected_last) + np.testing.assert_allclose(tactile, expected_tactile) + np.testing.assert_allclose( + env._prev_tactile_force, + np.array([[0.2, 0.01, 0.5, 0.3, 0.9]], dtype=np.float64), + ) + + +def test_sharpa_tactile_disabled_clears_history() -> None: + """Verify disabled tactile observation clears cached tactile history. + + Args: + None. + + Returns: + None. The assertions validate the disabled-tactile contract. + """ + sensor_data = { + "contact_right_thumb_elastomer_force": np.array([[1.0, 0.0, 0.0]], dtype=np.float64), + "contact_right_index_elastomer_force": np.array([[2.0, 0.0, 0.0]], dtype=np.float64), + "contact_right_middle_elastomer_force": np.array([[3.0, 0.0, 0.0]], dtype=np.float64), + "contact_right_ring_elastomer_force": np.array([[4.0, 0.0, 0.0]], dtype=np.float64), + "contact_right_pinky_elastomer_force": np.array([[5.0, 0.0, 0.0]], dtype=np.float64), + } + env = _build_fake_tactile_env( + sensor_data, + enable_tactile=False, + last_contacts=np.full((1, 5), 3.0, dtype=np.float64), + prev_tactile_force=np.full((1, 5), 4.0, dtype=np.float64), + ) + + tactile = SharpaInhandBaseEnv._compute_tactile_observation(env) + + np.testing.assert_allclose(tactile, np.zeros((1, 5), dtype=np.float64)) + np.testing.assert_allclose(env.last_contacts, np.zeros((1, 5), dtype=np.float64)) + np.testing.assert_allclose(env._prev_tactile_force, np.zeros((1, 5), dtype=np.float64)) + + +@pytest.mark.slow +def test_sharpa_mujoco_reset_applies_friction_randomization() -> None: + _require_mujoco_runtime() + ensure_registries() + + from unilab.base import registry + + num_envs = 4 + cfg, env_cfg_override = _compose_sharpa_mujoco_owner_cfg(num_envs) + with TemporaryDirectory() as tmp_dir: + cache_prefix = Path(tmp_dir) / "sharpa_grasp" + env_cfg_override["grasp_cache_path"] = str(cache_prefix) + for scale_value in env_cfg_override["domain_rand"]["scale_list"]: + cache_file = cache_prefix.parent / f"{cache_prefix.name}_{float(scale_value):g}.npy" + np.save(cache_file, np.zeros((8, 29), dtype=np.float32)) + + env = registry.make( + "SharpaInhandRotation", + num_envs=num_envs, + sim_backend="mujoco", + env_cfg_override=env_cfg_override, + ) + env_obj: Any = env + try: + env_ids = np.arange(num_envs, dtype=np.int32) + _, info = env_obj.reset(env_ids) + + backend: Any = env_obj._backend + pool = backend._pool + geom_friction = np.stack( + [pool.get_field(i, "geom_friction") for i in range(num_envs)], + axis=0, + ).reshape(num_envs, backend.model.ngeom, 3) + friction_scale = np.asarray(info["critic_info"][:, 3], dtype=np.float64) + + assert np.unique(np.round(friction_scale, 6)).size > 1 + assert np.all(friction_scale >= cfg.env.domain_rand.randomize_friction_scale_lower) + assert np.all(friction_scale <= cfg.env.domain_rand.randomize_friction_scale_upper) + + for env_idx in range(num_envs): + scale = friction_scale[env_idx] + for material, base_friction in ( + ("object", cfg.env.domain_rand.object_base_friction), + ("metal", cfg.env.domain_rand.metal_base_friction), + ("elastomer", cfg.env.domain_rand.elastomer_base_friction), + ): + actual = geom_friction[env_idx, env_obj._friction_geom_ids[material]] + expected = env_obj._friction_profile(material, base_friction) * scale + np.testing.assert_allclose(actual, np.broadcast_to(expected, actual.shape)) + finally: + pool = getattr(getattr(env_obj, "_backend", None), "_pool", None) + if pool is not None: + pool.close() + env_obj.close() + + +@pytest.mark.slow +def test_sharpa_mujoco_reset_randomizes_pd_gains_from_xml_defaults() -> None: + """Verify Sharpa reset PD gains scale MuJoCo XML actuator defaults per DOF. + + Args: + None. + + Returns: + None. The assertions validate info buffers and backend reset payloads. + """ + _require_mujoco_runtime() + ensure_registries() + + from unilab.base import registry + + num_envs = 4 + cfg, env_cfg_override = _compose_sharpa_mujoco_owner_cfg(num_envs) + with TemporaryDirectory() as tmp_dir: + cache_prefix = Path(tmp_dir) / "sharpa_grasp" + env_cfg_override["grasp_cache_path"] = str(cache_prefix) + for scale_value in env_cfg_override["domain_rand"]["scale_list"]: + cache_file = cache_prefix.parent / f"{cache_prefix.name}_{float(scale_value):g}.npy" + np.save(cache_file, np.zeros((8, 29), dtype=np.float32)) + + env = registry.make( + "SharpaInhandRotation", + num_envs=num_envs, + sim_backend="mujoco", + env_cfg_override=env_cfg_override, + ) + env_obj: Any = env + try: + env_ids = np.arange(num_envs, dtype=np.int32) + _, info = env_obj.reset(env_ids) + + backend: Any = env_obj._backend + pool = backend._pool + default_kp, default_kd = backend.get_actuator_gains() + default_kp = np.asarray(default_kp[: env_obj._num_action], dtype=np.float64) + default_kd = np.asarray(default_kd[: env_obj._num_action], dtype=np.float64) + info_kp = np.asarray(info["p_gain"], dtype=np.float64) + info_kd = np.asarray(info["d_gain"], dtype=np.float64) + pool_kp = np.stack([pool.get_field(i, "kp") for i in range(num_envs)], axis=0) + pool_kd = np.stack([pool.get_field(i, "kd") for i in range(num_envs)], axis=0) + + assert np.all(default_kp > 0.0) + assert np.all(default_kd > 0.0) + assert info_kp.shape == (num_envs, env_obj._num_action) + assert info_kd.shape == (num_envs, env_obj._num_action) + np.testing.assert_allclose(pool_kp[:, : env_obj._num_action], info_kp) + np.testing.assert_allclose(pool_kd[:, : env_obj._num_action], info_kd) + + kp_scale = info_kp / default_kp[None, :] + kd_scale = info_kd / default_kd[None, :] + assert np.unique(np.round(kp_scale.reshape(-1), 6)).size > 1 + assert np.unique(np.round(kd_scale.reshape(-1), 6)).size > 1 + assert np.all(kp_scale >= cfg.env.domain_rand.randomize_p_gain_scale_lower) + assert np.all(kp_scale <= cfg.env.domain_rand.randomize_p_gain_scale_upper) + assert np.all(kd_scale >= cfg.env.domain_rand.randomize_d_gain_scale_lower) + assert np.all(kd_scale <= cfg.env.domain_rand.randomize_d_gain_scale_upper) + finally: + pool = getattr(getattr(env_obj, "_backend", None), "_pool", None) + if pool is not None: + pool.close() + env_obj.close() + + +@pytest.mark.slow +def test_sharpa_mujoco_interval_force_disturbs_object_velocity() -> None: + """Verify Sharpa force randomization perturbs object linear velocity through DR. + + Args: + None. + + Returns: + None. The assertions validate the interval-randomization force path. + """ + _require_mujoco_runtime() + ensure_registries() + + from unilab.base import registry + + num_envs = 4 + _, env_cfg_override = _compose_sharpa_mujoco_owner_cfg(num_envs) + env_cfg_override["domain_rand"]["force_scale"] = 2.0 + env_cfg_override["domain_rand"]["random_force_prob_scalar"] = 1.0 + with TemporaryDirectory() as tmp_dir: + cache_prefix = Path(tmp_dir) / "sharpa_grasp" + env_cfg_override["grasp_cache_path"] = str(cache_prefix) + for scale_value in env_cfg_override["domain_rand"]["scale_list"]: + cache_file = cache_prefix.parent / f"{cache_prefix.name}_{float(scale_value):g}.npy" + np.save(cache_file, np.zeros((8, 29), dtype=np.float32)) + + env = registry.make( + "SharpaInhandRotation", + num_envs=num_envs, + sim_backend="mujoco", + env_cfg_override=env_cfg_override, + ) + env_obj: Any = env + try: + env.init_state() + assert env_obj._backend.get_dr_capabilities().supports_interval_body_force + before = env_obj._backend._physics_state.copy() + env_obj._dr_manager.apply_interval_randomization_if_due(env_obj.step_counter) + + body_id = int(env_obj._object_body_id) + force_slice = slice(6 * body_id, 6 * body_id + 3) + joint_adr = int(env_obj._backend.model.body_jntadr[body_id]) + dof_adr = int(env_obj._backend.model.jnt_dofadr[joint_adr]) + qvel_slice = slice( + env_obj._backend._idx_qvel + dof_adr, + env_obj._backend._idx_qvel + dof_adr + 3, + ) + + assert np.any(np.linalg.norm(env_obj._random_object_force, axis=1) > 0.0) + np.testing.assert_allclose( + env_obj._backend._pending_xfrc_applied[:, force_slice], + env_obj._random_object_force, + ) + env_obj._backend.step( + np.zeros((num_envs, env_obj._num_action), dtype=env_obj._np_dtype), + nsteps=1, + ) + after = env_obj._backend._physics_state + velocity_delta = after[:, qvel_slice] - before[:, qvel_slice] + assert np.any(np.linalg.norm(velocity_delta, axis=1) > 0.0) + finally: + pool = getattr(getattr(env_obj, "_backend", None), "_pool", None) + if pool is not None: + pool.close() + env_obj.close() + + +@pytest.mark.slow +def test_sharpa_mujoco_interval_force_plan_matches_decay_and_mass_scaled_resample( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Verify Sharpa force DR matches decay plus mass-scaled Gaussian resampling. + + Args: + monkeypatch: Pytest helper used to make numpy random sampling deterministic. + + Returns: + None. The assertions validate the exact interval body-force payload. + """ + _require_mujoco_runtime() + ensure_registries() + + from unilab.base import registry + + num_envs = 4 + _, env_cfg_override = _compose_sharpa_mujoco_owner_cfg(num_envs) + env_cfg_override["domain_rand"]["force_scale"] = 2.0 + env_cfg_override["domain_rand"]["random_force_prob_scalar"] = 0.5 + env_cfg_override["domain_rand"]["randomize_mass"] = False + with TemporaryDirectory() as tmp_dir: + cache_prefix = Path(tmp_dir) / "sharpa_grasp" + env_cfg_override["grasp_cache_path"] = str(cache_prefix) + for scale_value in env_cfg_override["domain_rand"]["scale_list"]: + cache_file = cache_prefix.parent / f"{cache_prefix.name}_{float(scale_value):g}.npy" + np.save(cache_file, np.zeros((8, 29), dtype=np.float32)) + + env = registry.make( + "SharpaInhandRotation", + num_envs=num_envs, + sim_backend="mujoco", + env_cfg_override=env_cfg_override, + ) + env_obj: Any = env + try: + env.init_state() + env_obj._random_object_force[:] = np.array( + [ + [0.4, -0.2, 0.1], + [-0.3, 0.5, -0.7], + [0.8, -0.6, 0.2], + [0.1, 0.3, -0.4], + ], + dtype=np.float64, + ) + previous_force = env_obj._random_object_force.copy() + sampled_uniform = np.array([0.1, 0.9, 0.2, 0.8], dtype=np.float64) + sampled_gaussian = np.array( + [ + [1.0, -2.0, 0.5], + [-1.5, 0.25, 2.0], + ], + dtype=np.float64, + ) + + def _fake_rand(*shape: int) -> np.ndarray: + assert shape == (num_envs,) + return sampled_uniform.copy() + + def _fake_randn(*shape: int) -> np.ndarray: + assert shape == (2, 3) + return sampled_gaussian.copy() + + monkeypatch.setattr(np.random, "rand", _fake_rand) + monkeypatch.setattr(np.random, "randn", _fake_randn) + + decay = float( + np.power( + env_obj.cfg.domain_rand.force_decay, + env_obj.cfg.ctrl_dt / max(env_obj.cfg.domain_rand.force_decay_interval, 1.0e-8), + ) + ) + object_mass = env_obj._resolve_current_object_mass() + resample_mask = sampled_uniform < float( + env_obj.cfg.domain_rand.random_force_prob_scalar + ) + + plan = SharpaInhandRotationDRProvider().build_interval_randomization_plan( + env_obj, + step_counter=0, + ) + + assert plan is not None + assert plan.body_ids is not None + assert plan.body_force is not None + np.testing.assert_array_equal( + plan.body_ids, + np.asarray([env_obj._object_body_id], dtype=np.int32), + ) + + # Envs below the Bernoulli threshold get a fresh Gaussian force sample; + # the remaining envs keep the decayed previous force. + expected_force = previous_force * decay + expected_force[resample_mask] = ( + sampled_gaussian + * object_mass[resample_mask, None] + * float(env_obj.cfg.domain_rand.force_scale) + ) + np.testing.assert_allclose(env_obj._random_object_force, expected_force) + np.testing.assert_allclose(plan.body_force, expected_force[:, None, :]) + assert plan.body_linear_velocity_delta is None + finally: + pool = getattr(getattr(env_obj, "_backend", None), "_pool", None) + if pool is not None: + pool.close() + env_obj.close() + + +@pytest.mark.slow +def test_sharpa_mujoco_reset_applies_object_mass_and_com_randomization() -> None: + _require_mujoco_runtime() + ensure_registries() + + from unilab.base import registry + + num_envs = 4 + cfg, env_cfg_override = _compose_sharpa_mujoco_owner_cfg(num_envs) + with TemporaryDirectory() as tmp_dir: + cache_prefix = Path(tmp_dir) / "sharpa_grasp" + env_cfg_override["grasp_cache_path"] = str(cache_prefix) + for scale_value in env_cfg_override["domain_rand"]["scale_list"]: + cache_file = cache_prefix.parent / f"{cache_prefix.name}_{float(scale_value):g}.npy" + np.save(cache_file, np.zeros((8, 29), dtype=np.float32)) + + env = registry.make( + "SharpaInhandRotation", + num_envs=num_envs, + sim_backend="mujoco", + env_cfg_override=env_cfg_override, + ) + env_obj: Any = env + try: + env_ids = np.arange(num_envs, dtype=np.int32) + _, info = env_obj.reset(env_ids) + + backend: Any = env_obj._backend + pool = backend._pool + body_mass = np.stack([pool.get_field(i, "body_mass") for i in range(num_envs)], axis=0) + body_ipos = np.stack([pool.get_field(i, "body_ipos") for i in range(num_envs)], axis=0) + body_ipos = body_ipos.reshape(num_envs, backend.model.nbody, 3) + + object_body_id = int(env_obj._object_body_id) + randomized_mass = np.asarray(info["critic_info"][:, 4], dtype=np.float64) + randomized_com = np.asarray(info["critic_info"][:, 5:8], dtype=np.float64) + + assert np.unique(np.round(randomized_mass, 6)).size > 1 + assert np.unique(np.round(randomized_com.reshape(-1), 6)).size > 1 + assert np.all(randomized_mass >= cfg.env.domain_rand.randomize_mass_lower) + assert np.all(randomized_mass <= cfg.env.domain_rand.randomize_mass_upper) + assert np.all(randomized_com >= cfg.env.domain_rand.randomize_com_lower) + assert np.all(randomized_com <= cfg.env.domain_rand.randomize_com_upper) + + np.testing.assert_allclose(body_mass[:, object_body_id], randomized_mass) + np.testing.assert_allclose( + body_ipos[:, object_body_id, :], + env_obj._base_body_ipos[object_body_id][None, :] + randomized_com, + ) + finally: + pool = getattr(getattr(env_obj, "_backend", None), "_pool", None) + if pool is not None: + pool.close() + env_obj.close() diff --git a/tests/envs/test_sharpa_geom_scale.py b/tests/envs/test_sharpa_geom_scale.py deleted file mode 100644 index 769350204..000000000 --- a/tests/envs/test_sharpa_geom_scale.py +++ /dev/null @@ -1,49 +0,0 @@ -from __future__ import annotations - -from types import SimpleNamespace - -import numpy as np - -from unilab.envs.manipulation.sharpa_inhand.rotation import SharpaInhandRotationDRProvider - - -def test_sharpa_provider_builds_mujoco_init_geom_scale_plan() -> None: - env = SimpleNamespace( - _backend=SimpleNamespace(backend_type="mujoco"), - _object_geom_base_size=np.array([0.02, 0.016, 0.0], dtype=np.float64), - scale_values=np.array([0.5, 0.8], dtype=np.float64), - scale_ids=np.array([0, 0, 1, 1], dtype=np.int32), - cfg=SimpleNamespace(object_geom_name="object"), - ) - - plan = SharpaInhandRotationDRProvider().build_init_randomization_plan(env) - - assert plan is not None - np.testing.assert_array_equal( - plan.model_assignments, - np.array([0, 0, 1, 1], dtype=np.int32), - ) - assert len(plan.model_variants) == 2 - assert plan.model_variants[0].geom_size_overrides[0].geom_name == "object" - np.testing.assert_allclose( - plan.model_variants[0].geom_size_overrides[0].size, - [0.01, 0.008, 0.0], - ) - np.testing.assert_allclose( - plan.model_variants[1].geom_size_overrides[0].size, - [0.016, 0.0128, 0.0], - ) - - -def test_sharpa_provider_skips_non_mujoco_init_geom_scale_plan() -> None: - env = SimpleNamespace( - _backend=SimpleNamespace(backend_type="motrix"), - _object_geom_base_size=np.array([0.02, 0.016, 0.0], dtype=np.float64), - scale_values=np.array([0.5], dtype=np.float64), - scale_ids=np.array([0], dtype=np.int32), - cfg=SimpleNamespace(object_geom_name="object"), - ) - - plan = SharpaInhandRotationDRProvider().build_init_randomization_plan(env) - - assert plan is None diff --git a/tests/ipc/test_shared_onpolicy_storage.py b/tests/ipc/test_shared_onpolicy_storage.py index 7ab8b627c..00316b0f8 100644 --- a/tests/ipc/test_shared_onpolicy_storage.py +++ b/tests/ipc/test_shared_onpolicy_storage.py @@ -196,3 +196,10 @@ def test_storage_allocates_optional_critic_fields(): assert set(storage.write_buffer.keys()) == expected storage.cleanup() + + +def test_storage_ipc_contract_has_no_privileged_fields(): + from unilab.ipc.shared_onpolicy_storage import _FIELD_SHAPES + + assert "priv_info" not in _FIELD_SHAPES + assert "last_priv_info" not in _FIELD_SHAPES diff --git a/tests/scripts/doc_checks.py b/tests/scripts/doc_checks.py index 2af115836..914f70cb1 100644 --- a/tests/scripts/doc_checks.py +++ b/tests/scripts/doc_checks.py @@ -11,6 +11,7 @@ "training", "reward", "env", + "teacher", "interactive", "viser", "num_envs", diff --git a/tests/scripts/test_check_docs.py b/tests/scripts/test_check_docs.py index 50f1015ee..c2a205d20 100644 --- a/tests/scripts/test_check_docs.py +++ b/tests/scripts/test_check_docs.py @@ -2,6 +2,8 @@ from pathlib import Path +from omegaconf import OmegaConf + from tests.scripts import doc_checks @@ -11,6 +13,25 @@ def test_documentation_files_match_current_repo_contracts(): assert errors == [] +def test_sharpa_domain_randomization_doc_matches_owner_config(): + root = Path(__file__).resolve().parents[2] + doc_path = root / "docs" / "users" / "zh_CN" / "06-domain-randomization.md" + content = doc_path.read_text(encoding="utf-8") + owner_cfg = OmegaConf.load(root / "conf" / "ppo" / "task" / "sharpa_inhand" / "mujoco.yaml") + + sharpa_rows = [ + line for line in content.splitlines() if line.startswith("| `SharpaInhandRotation` |") + ] + assert sharpa_rows + if float(owner_cfg.env.domain_rand.force_scale) > 0.0: + assert all("| 无 |" not in row for row in sharpa_rows) + assert any("body_force" in row for row in sharpa_rows) + + assert "`env.domain_rand.scale_list`" in content + assert "`env.scale_list`" not in content + assert "必须能被 `num_scales` 整除" not in content + + def test_check_training_entrypoint_semantics_flags_issue_204_patterns(): root = Path(__file__).resolve().parents[2] doc_path = root / "README.md" diff --git a/tests/scripts/test_support_matrix.py b/tests/scripts/test_support_matrix.py index 57b60b8a4..19f9e344f 100644 --- a/tests/scripts/test_support_matrix.py +++ b/tests/scripts/test_support_matrix.py @@ -32,3 +32,19 @@ def test_support_matrix_keeps_uncovered_mlx_tasks_at_configured(): assert row.cells["mujoco"].level == EvidenceLevel.CONFIGURED assert row.cells["motrix"].level == EvidenceLevel.CONFIGURED + + +def test_support_matrix_marks_sharpa_motrix_as_missing(): + row = _row("PPO (torch)", "sharpa_inhand") + + assert row.cells["mujoco"].level == EvidenceLevel.TESTED + assert row.cells["motrix"].level == EvidenceLevel.MISSING + + appo_row = _row("APPO (torch)", "sharpa_inhand") + + assert appo_row.cells["mujoco"].level == EvidenceLevel.TESTED + assert appo_row.cells["motrix"].level == EvidenceLevel.MISSING + allegro_appo_row = _row("APPO (torch)", "allegro_inhand") + + assert allegro_appo_row.cells["mujoco"].level == EvidenceLevel.TESTED + assert allegro_appo_row.cells["motrix"].level == EvidenceLevel.TESTED diff --git a/tests/scripts/test_train_scripts.py b/tests/scripts/test_train_scripts.py index bc35ffd6e..50dbac616 100644 --- a/tests/scripts/test_train_scripts.py +++ b/tests/scripts/test_train_scripts.py @@ -19,6 +19,7 @@ import pytest from hydra import compose, initialize_config_dir from hydra.core.global_hydra import GlobalHydra +from omegaconf import OmegaConf _SCRIPTS_DIR = Path(__file__).parent.parent.parent / "scripts" _CONF_DIR = Path(__file__).parent.parent.parent / "conf" @@ -112,6 +113,20 @@ def _appo_cfg(overrides=None): return compose("config", overrides=_normalize_overrides(overrides)) +def _hora_distill_cfg(overrides=None): + """Compose the HORA distillation Hydra config. + + Args: + overrides: Optional Hydra override strings to apply during composition. + + Returns: + The composed HORA distillation config. + """ + GlobalHydra.instance().clear() + with initialize_config_dir(config_dir=str(_CONF_DIR / "hora_distill"), version_base="1.3"): + return compose("config", overrides=overrides or []) + + def _train_rsl_rl(monkeypatch: pytest.MonkeyPatch): import types @@ -132,6 +147,18 @@ def _train_appo(): return _load_script("train_appo") +def _train_hora_distill(): + """Load the HORA distillation entrypoint module. + + Args: + None. + + Returns: + The loaded ``scripts/train_hora_distill.py`` module. + """ + return _load_script("train_hora_distill") + + def test_offpolicy_hydra_default_algo(): cfg = _offpolicy_cfg() assert cfg.algo.algo == "sac" @@ -188,6 +215,70 @@ def test_offpolicy_hydra_algo_td3(): assert cfg.algo.algo == "td3" +def test_hora_distill_task_owner_overrides_root_config_defaults(): + mod = _train_hora_distill() + root_cfg = OmegaConf.load(_CONF_DIR / "hora_distill" / "config.yaml") + cfg = mod._apply_teacher_defaults(_hora_distill_cfg(["task=sharpa_inhand/mujoco"])) + + assert root_cfg.algo.num_envs == 4096 + assert root_cfg.algo.save_interval_steps == 100000000 + assert cfg.algo.num_envs == 16384 + assert cfg.algo.save_interval_steps == 10000000 + + +def test_hora_distill_script_delegates_teacher_owner_resolution(): + source = (_SCRIPTS_DIR / "train_hora_distill.py").read_text(encoding="utf-8") + + assert "OmegaConf.load" not in source + assert "HoraActorModel" not in source + assert 'conf" / str(algo_family)' not in source + + +@pytest.mark.parametrize("teacher_algo_family", ["ppo", "appo"]) +def test_hora_distill_teacher_owner_defaults_support_ppo_and_appo( + teacher_algo_family: str, +): + mod = _train_hora_distill() + cfg = mod._apply_teacher_defaults( + _hora_distill_cfg( + [ + "task=sharpa_inhand/mujoco", + f"teacher.algo_family={teacher_algo_family}", + "teacher.task=sharpa_inhand/mujoco_hora", + ] + ) + ) + + assert cfg.training.task_name == "SharpaInhandRotation" + assert cfg.training.sim_backend == "mujoco" + assert cfg.algo.model.priv_info_embed_dim == 9 + assert cfg.algo.model.priv_mlp_hidden_dims == [256, 128, 9] + + +@pytest.mark.parametrize("teacher_algo_family", ["ppo", "appo"]) +def test_hora_distill_teacher_run_slug_omits_teacher_run_name(teacher_algo_family: str): + mod = _train_hora_distill() + cfg = OmegaConf.create({"teacher": {"task": "sharpa_inhand/mujoco"}}) + teacher_checkpoint = Path("/tmp") / "2026-04-22_13-26-45_mujoco" / "model_10000.pt" + + metadata = mod._teacher_run_metadata( + cfg, + teacher_algo_family=teacher_algo_family, + teacher_checkpoint=teacher_checkpoint, + ) + + assert metadata["run_name"] == "2026-04-22_13-26-45_mujoco" + assert metadata["run_slug"] == f"teacher-{teacher_algo_family}" + + +def test_offpolicy_go1_motrix_task_is_not_configured(): + """SAC has no Go1 Motrix owner config; use PPO for Go1 joystick tasks.""" + from hydra.errors import MissingConfigException + + with pytest.raises(MissingConfigException, match="task/sac/go1_joystick_flat/motrix"): + _offpolicy_cfg(["task=sac/go1_joystick_flat/motrix"]) + + def test_offpolicy_g1_walk_flat_motrix_resolved_algo_matches_task_owner(): """Motrix SAC G1 walk flat composes backend-owned algo hyperparameters.""" cfg = _offpolicy_cfg(["task=sac/g1_walk_flat/motrix"]) @@ -345,15 +436,70 @@ def test_build_ppo_env_cfg_override_allegro_mujoco( ): mod = _train_rsl_rl(monkeypatch) cfg = _ppo_cfg(["task=allegro_inhand/mujoco"]) + ppo_motrix_cfg = _ppo_cfg(["task=allegro_inhand/motrix"]) + appo_cfg = _appo_cfg(["task=allegro_inhand/mujoco"]) + appo_motrix_cfg = _appo_cfg(["task=allegro_inhand/motrix"]) env_cfg_override = mod.build_ppo_env_cfg_override(cfg) assert cfg.training.task_name == "AllegroInhandRotation" + assert cfg.algo.empirical_normalization is False + assert cfg.algo.actor.obs_normalization is True + assert cfg.algo.critic.obs_normalization is True assert env_cfg_override["reward_config"]["scales"]["rotate"] == pytest.approx(1.25) assert env_cfg_override["reward_config"]["reset_z_threshold"] == pytest.approx(0.125) assert env_cfg_override["gen_grasp"] is False assert env_cfg_override["max_episode_seconds"] == pytest.approx(20.0) assert env_cfg_override["grasp_cache_path"] == "cache/allegro_grasp_50k.npy" + assert env_cfg_override["domain_rand"]["randomize_base_mass"] is False + assert env_cfg_override["domain_rand"]["random_com"] is False + assert env_cfg_override["domain_rand"]["randomize_gravity"] is False + assert env_cfg_override["domain_rand"]["push_robots"] is False + assert env_cfg_override["domain_rand"]["joint_noise"] == pytest.approx(0.0) + assert env_cfg_override["domain_rand"]["ball_vel_noise"] == pytest.approx(0.0) + assert env_cfg_override["domain_rand"]["ball_z_offset"] == pytest.approx(0.0) + assert appo_cfg.algo.num_envs == cfg.algo.num_envs + assert appo_cfg.algo.steps_per_env == cfg.algo.num_steps_per_env + assert appo_cfg.algo.max_iterations == cfg.algo.max_iterations + assert appo_cfg.algo.save_interval == cfg.algo.save_interval + assert list(appo_cfg.algo.actor.hidden_dims) == list(cfg.algo.actor.hidden_dims) + assert appo_cfg.algo.actor.activation == cfg.algo.actor.activation + assert appo_cfg.algo.actor.obs_normalization is True + assert list(appo_cfg.algo.critic.hidden_dims) == list(cfg.algo.critic.hidden_dims) + assert appo_cfg.algo.critic.activation == cfg.algo.critic.activation + assert appo_cfg.algo.critic.obs_normalization is True + assert appo_cfg.algo.algorithm.value_loss_coef == pytest.approx( + cfg.algo.algorithm.value_loss_coef + ) + assert appo_cfg.algo.algorithm.entropy_coef == pytest.approx(cfg.algo.algorithm.entropy_coef) + assert appo_cfg.algo.algorithm.learning_rate == pytest.approx(cfg.algo.algorithm.learning_rate) + assert appo_cfg.algo.algorithm.desired_kl == pytest.approx(cfg.algo.algorithm.desired_kl) + assert appo_cfg.algo.algorithm.num_learning_epochs == cfg.algo.algorithm.num_learning_epochs + assert appo_cfg.algo.algorithm.num_mini_batches == cfg.algo.algorithm.num_mini_batches + assert appo_cfg.algo.algorithm.clip_param == pytest.approx(cfg.algo.algorithm.clip_param) + assert appo_cfg.algo.algorithm.gamma == pytest.approx(cfg.algo.algorithm.gamma) + assert appo_cfg.algo.algorithm.lam == pytest.approx(cfg.algo.algorithm.lam) + assert appo_cfg.algo.algorithm.max_grad_norm == pytest.approx(cfg.algo.algorithm.max_grad_norm) + assert ( + appo_cfg.algo.algorithm.use_clipped_value_loss is cfg.algo.algorithm.use_clipped_value_loss + ) + assert appo_cfg.algo.algorithm.schedule == cfg.algo.algorithm.schedule + assert appo_motrix_cfg.training.task_name == appo_cfg.training.task_name + assert appo_motrix_cfg.training.sim_backend == ppo_motrix_cfg.training.sim_backend + assert appo_motrix_cfg.algo.num_envs == appo_cfg.algo.num_envs + assert appo_motrix_cfg.algo.steps_per_env == appo_cfg.algo.steps_per_env + assert appo_motrix_cfg.algo.max_iterations == appo_cfg.algo.max_iterations + assert appo_motrix_cfg.algo.save_interval == appo_cfg.algo.save_interval + assert appo_motrix_cfg.algo.actor.obs_normalization is True + assert appo_motrix_cfg.algo.critic.obs_normalization is True + assert appo_motrix_cfg.reward.scales.rotate == pytest.approx( + ppo_motrix_cfg.reward.scales.rotate + ) + assert appo_motrix_cfg.env.gen_grasp is ppo_motrix_cfg.env.gen_grasp + assert appo_motrix_cfg.env.domain_rand.randomize_base_mass is False + assert appo_motrix_cfg.env.domain_rand.random_com is False + assert appo_motrix_cfg.env.domain_rand.randomize_gravity is False + assert appo_motrix_cfg.env.domain_rand.push_robots is False def test_build_ppo_env_cfg_override_allegro_grasp_mujoco( @@ -365,13 +511,19 @@ def test_build_ppo_env_cfg_override_allegro_grasp_mujoco( env_cfg_override = mod.build_ppo_env_cfg_override(cfg) assert cfg.training.task_name == "AllegroInhandRotationGrasp" + assert cfg.algo.empirical_normalization is False + assert cfg.algo.actor.obs_normalization is True + assert cfg.algo.critic.obs_normalization is True assert env_cfg_override["reward_config"]["scales"]["rotate"] == pytest.approx(0.0) assert env_cfg_override["gen_grasp"] is True assert env_cfg_override["grasp_collection_target"] == 50000 assert env_cfg_override["grasp_quality_check"] is True assert env_cfg_override["domain_rand"]["randomize_base_mass"] is False assert env_cfg_override["domain_rand"]["random_com"] is False + assert env_cfg_override["domain_rand"]["randomize_gravity"] is False assert env_cfg_override["domain_rand"]["push_robots"] is False + assert env_cfg_override["domain_rand"]["ball_vel_noise"] == pytest.approx(0.0) + assert env_cfg_override["domain_rand"]["joint_noise"] == pytest.approx(0.25) def test_build_ppo_env_cfg_override_allegro_grasp_cli_override_wins( @@ -1169,6 +1321,59 @@ def reset(self, env_indices): assert wrapper.num_privileged_obs == 4 +def test_play_wrapper_preserves_hora_priv_info_and_proprio_history(): + import numpy as np + + from unilab.algos.torch.hora.rsl_rl import HoraRslRlVecEnvWrapper + + class FakeEnv: + def __init__(self): + self.num_envs = 1 + self.state = type( + "State", + (), + { + "obs": { + "obs": np.array([[1.0, 2.0, 3.0]], dtype=np.float32), + "critic": np.array([[1.0, 2.0, 3.0, 4.0, 5.0]], dtype=np.float32), + }, + "info": { + "critic_info": np.array([[4.0, 5.0]], dtype=np.float32), + "proprio_hist": np.array( + [[[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]], + dtype=np.float32, + ), + }, + }, + )() + self.cfg = type("Cfg", (), {"max_episode_seconds": 10.0, "ctrl_dt": 0.02})() + self.observation_space = type("Space", (), {"shape": (5,)})() + self.action_space = type("Space", (), {"shape": (2,)})() + self.obs_groups_spec = {"obs": 3, "critic": 5} + + def init_state(self): + pass + + def reset(self, env_indices): + del env_indices + return ( + cast(dict[str, np.ndarray], getattr(self.state, "obs")), + cast(dict[str, np.ndarray], getattr(self.state, "info")), + ) + + wrapper = HoraRslRlVecEnvWrapper(FakeEnv(), device="cpu", policy_obs_mode="flat") + obs_td, _ = wrapper.reset() + + np.testing.assert_allclose( + obs_td["priv_info"].cpu().numpy(), + np.array([[4.0, 5.0]], dtype=np.float32), + ) + np.testing.assert_allclose( + obs_td["proprio_hist"].cpu().numpy(), + np.array([[[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]], dtype=np.float32), + ) + + def test_play_wrapper_step_exports_timeout_bootstrap_obs(): import torch @@ -1226,6 +1431,77 @@ def step(self, actions): ) +def test_play_wrapper_timeout_bootstrap_preserves_hora_priv_info(): + import torch + + from unilab.algos.torch.hora.rsl_rl import HoraRslRlVecEnvWrapper + + class FakeEnv: + def __init__(self): + self.num_envs = 1 + self.cfg = type("Cfg", (), {"max_episode_seconds": 10.0, "ctrl_dt": 0.02})() + self.observation_space = type("Space", (), {"shape": (5,)})() + self.action_space = type("Space", (), {"shape": (2,)})() + self.obs_groups_spec = {"obs": 3, "critic": 5} + self.state = type( + "State", + (), + { + "obs": { + "obs": np.zeros((1, 3), dtype=np.float32), + "critic": np.zeros((1, 5), dtype=np.float32), + }, + "info": { + "critic_info": np.zeros((1, 2), dtype=np.float32), + "proprio_hist": np.zeros((1, 2, 3), dtype=np.float32), + }, + }, + )() + + def init_state(self): + pass + + def reset(self, env_indices): + del env_indices + return cast(dict[str, np.ndarray], getattr(self.state, "obs")), cast( + dict[str, np.ndarray], getattr(self.state, "info") + ) + + def step(self, actions): + del actions + return type( + "StepState", + (), + { + "obs": {"obs": np.array([[1.0, 2.0, 3.0]], dtype=np.float32)}, + "reward": np.array([1.0], dtype=np.float32), + "terminated": np.array([True]), + "truncated": np.array([True]), + "final_observation": { + "obs": np.array([[7.0, 8.0, 9.0]], dtype=np.float32), + "critic": np.array([[7.0, 8.0, 9.0, 4.0, 5.0]], dtype=np.float32), + }, + "info": { + "final_observation": { + "obs": np.array([[7.0, 8.0, 9.0]], dtype=np.float32), + "critic": np.array([[7.0, 8.0, 9.0, 4.0, 5.0]], dtype=np.float32), + }, + "critic_info": np.array([[0.0, 0.0]], dtype=np.float32), + "proprio_hist": np.zeros((1, 2, 3), dtype=np.float32), + }, + }, + )() + + wrapper = HoraRslRlVecEnvWrapper(FakeEnv(), device="cpu", policy_obs_mode="flat") + + _, _, _, infos = wrapper.step(torch.zeros((1, 2))) + + np.testing.assert_allclose( + infos["time_out_bootstrap_obs"]["priv_info"].cpu().numpy(), + np.array([[4.0, 5.0]], dtype=np.float32), + ) + + # --------------------------------------------------------------------------- # Issue #168: Unified log directory and load_run resolution # --------------------------------------------------------------------------- diff --git a/tests/test_sharpa.py b/tests/test_sharpa.py new file mode 100644 index 000000000..4f5059305 --- /dev/null +++ b/tests/test_sharpa.py @@ -0,0 +1,98 @@ +from __future__ import annotations + +from types import SimpleNamespace + +import numpy as np +import pytest + +from unilab.envs.common.rotation import np_quat_apply +from unilab.envs.manipulation.sharpa_inhand.base import SharpaDomainRandConfig +from unilab.envs.manipulation.sharpa_inhand.grasp_gen import ( + SharpaInhandRotationGraspCfg, + SharpaInhandRotationGraspEnv, +) +from unilab.envs.manipulation.sharpa_inhand.rotation import ( + SharpaInhandRotationEnv, + sample_random_quaternion, +) + + +def test_sharpa_gravity_direction_randomization_matches_rotated_gravity( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Gravity-direction DR should rotate a fixed-magnitude downward vector. + + Args: + monkeypatch: Pytest helper used to replace quaternion sampling. + + Returns: + None. The assertions validate the exact gravity vectors produced. + """ + fixed_quat = np.asarray( + [ + [1.0, 0.0, 0.0, 0.0], + [0.0, 0.0, 1.0, 0.0], + ], + dtype=np.float64, + ) + + def _fake_sample_random_quaternion(num_envs: int) -> np.ndarray: + assert num_envs == 2 + return fixed_quat.copy() + + monkeypatch.setitem( + SharpaInhandRotationEnv._build_gravity_direction_randomization.__globals__, + "sample_random_quaternion", + _fake_sample_random_quaternion, + ) + + env = SimpleNamespace( + _cfg=SimpleNamespace( + domain_rand=SimpleNamespace( + randomize_gravity_direction=True, + gravity_direction_magnitude=9.81, + ) + ) + ) + + gravity = SharpaInhandRotationEnv._build_gravity_direction_randomization(env, batch_size=2) + + assert gravity is not None + expected = np_quat_apply( + fixed_quat, + np.asarray([[0.0, 0.0, -9.81]], dtype=np.float64), + ) + np.testing.assert_allclose(gravity, expected) + np.testing.assert_allclose(np.linalg.norm(gravity, axis=1), np.full((2,), 9.81)) + + +def test_sharpa_gravity_direction_randomization_disabled_returns_none() -> None: + env = SimpleNamespace( + _cfg=SimpleNamespace( + domain_rand=SimpleNamespace( + randomize_gravity_direction=False, + gravity_direction_magnitude=9.81, + ) + ) + ) + + gravity = SharpaInhandRotationEnv._build_gravity_direction_randomization(env, batch_size=3) + + assert gravity is None + + +def test_sharpa_grasp_env_rejects_gravity_randomization() -> None: + """Sharpa grasp collection should reject gravity DR explicitly. + + Args: + None. + + Returns: + None. The assertion validates the grasp-task contract. + """ + cfg = SharpaInhandRotationGraspCfg( + domain_rand=SharpaDomainRandConfig(randomize_gravity_direction=True) + ) + + with pytest.raises(ValueError, match="does not support gravity randomization"): + SharpaInhandRotationGraspEnv(cfg, num_envs=1, backend_type="mujoco")