From 7b1c8e1507a00cf161c5dac633004ad556989e45 Mon Sep 17 00:00:00 2001 From: tatp-yf Date: Tue, 21 Apr 2026 19:58:23 +0800 Subject: [PATCH 1/4] refactor: unify g1 walk env tasks --- benchmark/benchmark_env_step.py | 4 +- .../benchmark_mujoco_backend_step_detail.py | 6 +- benchmark/benchmark_physics_step_genesis.py | 48 ++-- benchmark/benchmark_physics_step_isaacgym.py | 52 ++-- benchmark/benchmark_physics_step_mj_step.py | 81 +++--- .../benchmark_physics_step_mujoco_warp.py | 55 ++-- benchmark/core/task_names.py | 12 +- .../mujoco.yaml | 7 +- .../task/flashsac/g1_walk_flat/motrix.yaml | 9 + .../task/flashsac/g1_walk_flat/mujoco.yaml | 10 +- .../mujoco.yaml | 4 +- .../task/sac/g1_walk_flat/motrix.yaml | 8 + .../task/sac/g1_walk_flat/mujoco.yaml | 8 + .../task/sac/g1_walk_rough/mujoco.yaml | 8 + .../task/td3/g1_walk_flat/mujoco.yaml | 8 + .../motrix.yaml | 4 +- .../mujoco.yaml | 7 +- docs/users/zh_CN/02-simulation-backends.md | 8 +- docs/users/zh_CN/06-domain-randomization.md | 10 +- src/unilab/config/locomotion_params.py | 1 - src/unilab/docs/support_matrix.py | 14 +- src/unilab/envs/locomotion/g1/__init__.py | 10 +- src/unilab/envs/locomotion/g1/joystick.py | 228 +++++++++++++-- src/unilab/envs/locomotion/g1/joystick_sac.py | 242 ---------------- tests/algos/test_rsl_rl_runner.py | 2 +- tests/base/test_reward_override.py | 32 ++- tests/config/test_config_system.py | 8 +- tests/config/test_locomotion_params.py | 57 +++- tests/config/test_reward_injection.py | 2 +- .../locomotion/g1/test_issue175_regression.py | 265 ++++++++++++++++++ .../locomotion/g1/test_symmetry_contract.py | 6 +- tests/envs/test_env_configs.py | 133 +++++++-- tests/scripts/test_train_script_configs.py | 2 +- tests/scripts/test_train_scripts.py | 20 +- 34 files changed, 904 insertions(+), 467 deletions(-) rename conf/appo/task/{g1_joystick_flat => g1_walk_flat}/mujoco.yaml (86%) rename conf/offpolicy/task/flashsac/{g1_joystick_flat => g1_walk_flat_amp}/mujoco.yaml (94%) rename conf/ppo/task/{g1_joystick_flat => g1_walk_flat}/motrix.yaml (95%) rename conf/ppo/task/{g1_joystick_flat => g1_walk_flat}/mujoco.yaml (87%) delete mode 100644 src/unilab/envs/locomotion/g1/joystick_sac.py create mode 100644 tests/envs/locomotion/g1/test_issue175_regression.py diff --git a/benchmark/benchmark_env_step.py b/benchmark/benchmark_env_step.py index e63c9bdc8..a020ffd1b 100644 --- a/benchmark/benchmark_env_step.py +++ b/benchmark/benchmark_env_step.py @@ -5,7 +5,7 @@ uv run benchmark/benchmark_env_step.py # Single task + backend: - uv run benchmark/benchmark_env_step.py task=g1_joystick_flat/motrix + uv run benchmark/benchmark_env_step.py task=g1_walk_flat/motrix # Override bench params: uv run benchmark/benchmark_env_step.py task=go1_joystick_flat/mujoco num_envs=4096 num_steps=500 @@ -61,7 +61,7 @@ def _load_helper_module(module_name: str, relative_path: str): TASK_CONFIGS = { "go1": "task=go1_joystick_flat", "go2": "task=go2_joystick_flat", - "g1": "task=g1_joystick_flat", + "g1": "task=g1_walk_flat", "g1_mt": "task=g1_motion_tracking", } diff --git a/benchmark/benchmark_mujoco_backend_step_detail.py b/benchmark/benchmark_mujoco_backend_step_detail.py index 742327860..676274487 100644 --- a/benchmark/benchmark_mujoco_backend_step_detail.py +++ b/benchmark/benchmark_mujoco_backend_step_detail.py @@ -19,12 +19,12 @@ uv run benchmark/benchmark_mujoco_backend_step_detail.py uv run benchmark/benchmark_mujoco_backend_step_detail.py \ - --tasks go1_joystick_flat,go2_joystick_flat,g1_joystick_flat \ + --tasks go1_joystick_flat,go2_joystick_flat,g1_walk_flat \ --env-nums 256,512,1024,2048,4096,8192 \ --nsteps 1,2,3,4 uv run benchmark/benchmark_mujoco_backend_step_detail.py \ - --tasks g1_joystick_flat --env-nums 2048 --nsteps 1,2,3 --iters 20 + --tasks g1_walk_flat --env-nums 2048 --nsteps 1,2,3 --iters 20 """ from __future__ import annotations @@ -89,7 +89,7 @@ def _load_helper_module(module_name: str, relative_path: str): TASK_COLORS = { "go1_joystick_flat": "#4C78A8", "go2_joystick_flat": "#54A24B", - "g1_joystick_flat": "#F58518", + "g1_walk_flat": "#F58518", } COMPONENT_COLORS = { "set_ctrl_ms": "#9AA0A6", diff --git a/benchmark/benchmark_physics_step_genesis.py b/benchmark/benchmark_physics_step_genesis.py index 18261f42d..9274f8846 100644 --- a/benchmark/benchmark_physics_step_genesis.py +++ b/benchmark/benchmark_physics_step_genesis.py @@ -3,7 +3,7 @@ Benchmark Genesis physics execution. Benchmarks Genesis across current locomotion owner task ids -(go1_joystick_flat/go2_joystick_flat/g1_joystick_flat) and outputs JSON + plots +(go1_joystick_flat/go2_joystick_flat/g1_walk_flat) and outputs JSON + plots aligned with benchmark/benchmark_physics_step_mj_step.py. Legacy env names remain accepted as aliases. @@ -20,7 +20,7 @@ from dataclasses import asdict, dataclass from datetime import datetime, timezone from pathlib import Path -from typing import Dict, List +from typing import Any, Dict, List, cast import matplotlib @@ -28,24 +28,30 @@ import matplotlib.pyplot as plt try: - import genesis as gs + import genesis as _gs except ImportError: gs = None +else: + gs = _gs try: - from benchmark.core.device_info import get_device_info_dict, get_device_info_line - from benchmark.core.task_names import ( - canonical_locomotion_task_ids, - locomotion_task_spec, - normalize_locomotion_task_id, - ) + from benchmark.core import device_info as _benchmark_device_info + from benchmark.core import task_names as _benchmark_task_names + + _device_info = _benchmark_device_info + _task_names = _benchmark_task_names except ModuleNotFoundError: - from core.device_info import get_device_info_dict, get_device_info_line - from core.task_names import ( - canonical_locomotion_task_ids, - locomotion_task_spec, - normalize_locomotion_task_id, - ) + from core import device_info as _core_device_info + from core import task_names as _core_task_names + + _device_info = _core_device_info + _task_names = _core_task_names + +get_device_info_dict = _device_info.get_device_info_dict +get_device_info_line = _device_info.get_device_info_line +canonical_locomotion_task_ids = _task_names.canonical_locomotion_task_ids +locomotion_task_spec = _task_names.locomotion_task_spec +normalize_locomotion_task_id = _task_names.normalize_locomotion_task_id @dataclass @@ -84,7 +90,8 @@ def _init_genesis() -> None: global _GENESIS_INITIALIZED if _GENESIS_INITIALIZED: return - gs.init(backend=gs.gpu) + gs_mod = cast(Any, gs) + gs_mod.init(backend=gs_mod.gpu) _GENESIS_INITIALIZED = True @@ -94,14 +101,15 @@ def _load_task_xml(task_name: str) -> str: def _build_scene(xml_path: str, batch_size: int): - scene = gs.Scene( + gs_mod = cast(Any, gs) + scene = gs_mod.Scene( show_viewer=False, - rigid_options=gs.options.RigidOptions( + rigid_options=gs_mod.options.RigidOptions( dt=0.01, - constraint_solver=gs.constraint_solver.Newton, + constraint_solver=gs_mod.constraint_solver.Newton, ), ) - scene.add_entity(gs.morphs.MJCF(file=xml_path)) + scene.add_entity(gs_mod.morphs.MJCF(file=xml_path)) scene.build(n_envs=batch_size) return scene diff --git a/benchmark/benchmark_physics_step_isaacgym.py b/benchmark/benchmark_physics_step_isaacgym.py index ef81106cc..8f7cf03a5 100644 --- a/benchmark/benchmark_physics_step_isaacgym.py +++ b/benchmark/benchmark_physics_step_isaacgym.py @@ -25,7 +25,7 @@ from dataclasses import asdict, dataclass from datetime import datetime, timezone from pathlib import Path -from typing import Dict, List +from typing import Any, Dict, List, cast import matplotlib @@ -43,13 +43,13 @@ if str(DEFAULT_ISAACGYM_PYTHON) not in sys.path: sys.path.insert(0, str(DEFAULT_ISAACGYM_PYTHON)) +_ISAACGYM_IMPORT_ERROR: Exception | None = None + try: from isaacgym import gymapi except Exception as _isaacgym_error: gymapi = None _ISAACGYM_IMPORT_ERROR = _isaacgym_error -else: - _ISAACGYM_IMPORT_ERROR = None @dataclass(frozen=True) @@ -87,9 +87,9 @@ class BenchRecord: asset_file="go2_description/urdf/go2_description.urdf", initial_height=0.40, ), - "g1_joystick_flat": TaskSpec( - owner_task_id="g1_joystick_flat", - display_name="g1_joystick_flat", + "g1_walk_flat": TaskSpec( + owner_task_id="g1_walk_flat", + display_name="g1_walk_flat", asset_root=DEFAULT_MODELS_ROOT, asset_file="g1_description/g1_29dof_rev_1_0.urdf", initial_height=0.78, @@ -98,16 +98,16 @@ class BenchRecord: TASK_ALIASES = { "Go1JoystickFlat": "go1_joystick_flat", "Go2JoystickFlat": "go2_joystick_flat", - "G1JoystickFlat": "g1_joystick_flat", + "G1WalkFlat": "g1_walk_flat", "task=go1_joystick_flat/isaacgym": "go1_joystick_flat", "task=go2_joystick_flat/isaacgym": "go2_joystick_flat", - "task=g1_joystick_flat/isaacgym": "g1_joystick_flat", + "task=g1_walk_flat/isaacgym": "g1_walk_flat", "go1_joystick_flat/isaacgym": "go1_joystick_flat", "go2_joystick_flat/isaacgym": "go2_joystick_flat", - "g1_joystick_flat/isaacgym": "g1_joystick_flat", + "g1_walk_flat/isaacgym": "g1_walk_flat", "go1": "go1_joystick_flat", "go2": "go2_joystick_flat", - "g1": "g1_joystick_flat", + "g1": "g1_walk_flat", } DEFAULT_TASK_IDS = list(TASK_SPECS.keys()) DEFAULT_BATCH_SIZES = [2**k for k in range(8, 15)] # 256 .. 16384 @@ -196,19 +196,20 @@ def _default_dof_targets(dof_props) -> np.ndarray: def _create_sim(compute_device_id: int, graphics_device_id: int, nthread: int): - gym = gymapi.acquire_gym() - sim_params = gymapi.SimParams() + gymapi_mod = cast(Any, gymapi) + gym = gymapi_mod.acquire_gym() + sim_params = gymapi_mod.SimParams() sim_params.dt = 0.01 sim_params.substeps = 1 - sim_params.up_axis = gymapi.UpAxis.UP_AXIS_Z - sim_params.gravity = gymapi.Vec3(0.0, 0.0, -9.81) + sim_params.up_axis = gymapi_mod.UpAxis.UP_AXIS_Z + sim_params.gravity = gymapi_mod.Vec3(0.0, 0.0, -9.81) sim_params.physx.solver_type = 1 sim_params.physx.num_position_iterations = 4 sim_params.physx.num_velocity_iterations = 1 sim_params.physx.num_threads = nthread sim_params.physx.use_gpu = True sim_params.use_gpu_pipeline = True - sim = gym.create_sim(compute_device_id, graphics_device_id, gymapi.SIM_PHYSX, sim_params) + sim = gym.create_sim(compute_device_id, graphics_device_id, gymapi_mod.SIM_PHYSX, sim_params) if sim is None: raise RuntimeError("Failed to create Isaac Gym simulation.") return gym, sim @@ -223,31 +224,32 @@ def _build_task_sim( ): spec = _task_spec(task_name) gym, sim = _create_sim(compute_device_id, graphics_device_id, nthread) + gymapi_mod = cast(Any, gymapi) - plane_params = gymapi.PlaneParams() - plane_params.normal = gymapi.Vec3(0.0, 0.0, 1.0) + plane_params = gymapi_mod.PlaneParams() + plane_params.normal = gymapi_mod.Vec3(0.0, 0.0, 1.0) gym.add_ground(sim, plane_params) - asset_options = gymapi.AssetOptions() + asset_options = gymapi_mod.AssetOptions() asset_options.flip_visual_attachments = True asset_options.armature = 0.01 - asset_options.default_dof_drive_mode = int(gymapi.DOF_MODE_POS) + asset_options.default_dof_drive_mode = int(gymapi_mod.DOF_MODE_POS) asset = gym.load_asset(sim, str(spec.asset_root), spec.asset_file, asset_options) dof_props = gym.get_asset_dof_properties(asset) - dof_props["driveMode"][:].fill(gymapi.DOF_MODE_POS) + dof_props["driveMode"][:].fill(gymapi_mod.DOF_MODE_POS) dof_props["stiffness"][:].fill(1000.0) dof_props["damping"][:].fill(10.0) dof_targets = _default_dof_targets(dof_props) spacing = 1.0 - env_lower = gymapi.Vec3(-spacing, 0.0, -spacing) - env_upper = gymapi.Vec3(spacing, spacing, spacing) + env_lower = gymapi_mod.Vec3(-spacing, 0.0, -spacing) + env_upper = gymapi_mod.Vec3(spacing, spacing, spacing) num_per_row = max(1, int(math.sqrt(batch_size))) - pose = gymapi.Transform() - pose.r = gymapi.Quat(0.0, 0.0, 0.0, 1.0) - pose.p = gymapi.Vec3(0.0, 0.0, spec.initial_height) + pose = gymapi_mod.Transform() + pose.r = gymapi_mod.Quat(0.0, 0.0, 0.0, 1.0) + pose.p = gymapi_mod.Vec3(0.0, 0.0, spec.initial_height) for env_idx in range(batch_size): env = gym.create_env(sim, env_lower, env_upper, num_per_row) diff --git a/benchmark/benchmark_physics_step_mj_step.py b/benchmark/benchmark_physics_step_mj_step.py index cbf959e7a..1fd168909 100644 --- a/benchmark/benchmark_physics_step_mj_step.py +++ b/benchmark/benchmark_physics_step_mj_step.py @@ -6,13 +6,14 @@ Other: benchmarks mujoco.rollout with the configured thread count only. Sweeps batch sizes across current locomotion owner task ids -(go1_joystick_flat/go2_joystick_flat/g1_joystick_flat). +(go1_joystick_flat/go2_joystick_flat/g1_walk_flat). Legacy env names remain accepted as aliases. """ from __future__ import annotations import argparse +import importlib import json import platform import time @@ -20,7 +21,7 @@ from datetime import datetime, timezone from multiprocessing import cpu_count from pathlib import Path -from typing import Dict, List +from typing import Any, Dict, List, cast import matplotlib import mujoco @@ -32,32 +33,37 @@ from matplotlib.patches import Patch _IS_MACOS = platform.system() == "Darwin" +mx: Any = None +mj_mlx_step: Any = None if _IS_MACOS: - import mlx.core as mx + import mlx.core as _mx + + mx = _mx try: - from mujoco import mlx_step as mj_mlx_step + mj_mlx_step = importlib.import_module("mujoco.mlx_step") except Exception: mj_mlx_step = None -else: - mx = None - mj_mlx_step = None try: - from benchmark.core.device_info import get_device_info_dict, get_device_info_line - from benchmark.core.task_names import ( - canonical_locomotion_task_ids, - locomotion_task_spec, - normalize_locomotion_task_id, - ) + from benchmark.core import device_info as _benchmark_device_info + from benchmark.core import task_names as _benchmark_task_names + + _device_info = _benchmark_device_info + _task_names = _benchmark_task_names except ModuleNotFoundError: - from core.device_info import get_device_info_dict, get_device_info_line - from core.task_names import ( - canonical_locomotion_task_ids, - locomotion_task_spec, - normalize_locomotion_task_id, - ) + from core import device_info as _core_device_info + from core import task_names as _core_task_names + + _device_info = _core_device_info + _task_names = _core_task_names + +get_device_info_dict = _device_info.get_device_info_dict +get_device_info_line = _device_info.get_device_info_line +canonical_locomotion_task_ids = _task_names.canonical_locomotion_task_ids +locomotion_task_spec = _task_names.locomotion_task_spec +normalize_locomotion_task_id = _task_names.normalize_locomotion_task_id @dataclass @@ -73,19 +79,20 @@ class BenchRecord: DEFAULT_TASK_IDS = canonical_locomotion_task_ids() DEFAULT_BATCH_SIZES = [2**k for k in range(8, 15)] # 256 .. 16384 -TASK_ALPHA = {"go1_joystick_flat": 0.75, "go2_joystick_flat": 0.9, "g1_joystick_flat": 1.0} -TASK_HATCH = {"go1_joystick_flat": "//", "go2_joystick_flat": "\\\\", "g1_joystick_flat": "xx"} +TASK_ALPHA = {"go1_joystick_flat": 0.75, "go2_joystick_flat": 0.9, "g1_walk_flat": 1.0} +TASK_HATCH = {"go1_joystick_flat": "//", "go2_joystick_flat": "\\\\", "g1_walk_flat": "xx"} -def _keyframe0_state_and_ctrl(model: mujoco.MjModel) -> tuple[np.ndarray, np.ndarray]: - data = mujoco.MjData(model) +def _keyframe0_state_and_ctrl(model: Any) -> tuple[np.ndarray, np.ndarray]: + mujoco_mod = cast(Any, mujoco) + data = mujoco_mod.MjData(model) if model.nkey > 0: - mujoco.mj_resetDataKeyframe(model, data, 0) + mujoco_mod.mj_resetDataKeyframe(model, data, 0) else: - mujoco.mj_resetData(model, data) - nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS) + mujoco_mod.mj_resetData(model, data) + nstate = mujoco_mod.mj_stateSize(model, mujoco_mod.mjtState.mjSTATE_FULLPHYSICS) state0 = np.empty((nstate,), dtype=np.float64) - mujoco.mj_getState(model, data, state0, mujoco.mjtState.mjSTATE_FULLPHYSICS) + mujoco_mod.mj_getState(model, data, state0, mujoco_mod.mjtState.mjSTATE_FULLPHYSICS) if model.nu == 0: ctrl0 = np.empty((0,), dtype=np.float64) elif model.nkey > 0: @@ -139,10 +146,10 @@ def _run_mlx( control=control_mx, nstep=nstep, chunk_size=chunk_size, - out_dtype=mx.float32, + out_dtype=cast(Any, mx).float32, ) state_mx, sensor_mx = out if isinstance(out, tuple) else (out.state_mx, out.sensordata_mx) - mx.eval(state_mx, sensor_mx) + cast(Any, mx).eval(state_mx, sensor_mx) return (time.perf_counter() - t0) / niter @@ -150,9 +157,9 @@ def _has_native_mujoco_mlx_step() -> bool: return mj_mlx_step is not None and hasattr(mj_mlx_step, "MlxStepRunner") -def _load_task_model(task_name: str) -> mujoco.MjModel: +def _load_task_model(task_name: str) -> Any: cfg = locomotion_task_spec(task_name).config_cls() - return mujoco.MjModel.from_xml_path(cfg.model_file) + return cast(Any, mujoco).MjModel.from_xml_path(cfg.model_file) def _display_backend(backend: str) -> str: @@ -173,7 +180,7 @@ def _bench_one_task( task_key = normalize_locomotion_task_id(task_name) np.random.seed(42) model = _load_task_model(task_key) - nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS) + nstate = cast(Any, mujoco).mj_stateSize(model, cast(Any, mujoco).mjtState.mjSTATE_FULLPHYSICS) state0, ctrl0 = _keyframe0_state_and_ctrl(model) records: List[BenchRecord] = [] @@ -188,13 +195,13 @@ def _bench_one_task( if _IS_MACOS: actual_nthread = min(batch_size, nthread, cpu_count()) - data_list = [mujoco.MjData(model) for _ in range(actual_nthread)] + data_list = [cast(Any, mujoco).MjData(model) for _ in range(actual_nthread)] with ( mj_rollout.Rollout(nthread=actual_nthread) as numpy_runner, - mj_mlx_step.MlxStepRunner(nthread=actual_nthread) as mlx_runner, + cast(Any, mj_mlx_step).MlxStepRunner(nthread=actual_nthread) as mlx_runner, ): - initial_state_mx = mx.array(initial_state, dtype=mx.float32) - control_mx = mx.array(control, dtype=mx.float32) + initial_state_mx = cast(Any, mx).array(initial_state, dtype=cast(Any, mx).float32) + control_mx = cast(Any, mx).array(control, dtype=cast(Any, mx).float32) _run_numpy( numpy_runner, @@ -268,7 +275,7 @@ def _bench_one_task( f"mlx={mlx_t * 1000:.3f}ms ({batch_size * nstep / mlx_t / 1e4:.2f}万fps)" ) else: - data_list_n = [mujoco.MjData(model) for _ in range(nthread)] + data_list_n = [cast(Any, mujoco).MjData(model) for _ in range(nthread)] with mj_rollout.Rollout(nthread=nthread) as runner_n: _run_numpy( runner_n, diff --git a/benchmark/benchmark_physics_step_mujoco_warp.py b/benchmark/benchmark_physics_step_mujoco_warp.py index de73d704d..1222887d0 100644 --- a/benchmark/benchmark_physics_step_mujoco_warp.py +++ b/benchmark/benchmark_physics_step_mujoco_warp.py @@ -3,7 +3,7 @@ Benchmark MuJoCo Warp physics execution. Benchmarks mujoco_warp across current locomotion owner task ids -(go1_joystick_flat/go2_joystick_flat/g1_joystick_flat) and outputs JSON + plots +(go1_joystick_flat/go2_joystick_flat/g1_walk_flat) and outputs JSON + plots aligned with benchmark/benchmark_physics_step_mj_step.py. Legacy env names remain accepted as aliases. @@ -20,7 +20,7 @@ from dataclasses import asdict, dataclass from datetime import datetime, timezone from pathlib import Path -from typing import Dict, List +from typing import Any, Dict, List, cast import matplotlib @@ -44,19 +44,23 @@ warp = None try: - from benchmark.core.device_info import get_device_info_dict, get_device_info_line - from benchmark.core.task_names import ( - canonical_locomotion_task_ids, - locomotion_task_spec, - normalize_locomotion_task_id, - ) + from benchmark.core import device_info as _benchmark_device_info + from benchmark.core import task_names as _benchmark_task_names + + _device_info = _benchmark_device_info + _task_names = _benchmark_task_names except ModuleNotFoundError: - from core.device_info import get_device_info_dict, get_device_info_line - from core.task_names import ( - canonical_locomotion_task_ids, - locomotion_task_spec, - normalize_locomotion_task_id, - ) + from core import device_info as _core_device_info + from core import task_names as _core_task_names + + _device_info = _core_device_info + _task_names = _core_task_names + +get_device_info_dict = _device_info.get_device_info_dict +get_device_info_line = _device_info.get_device_info_line +canonical_locomotion_task_ids = _task_names.canonical_locomotion_task_ids +locomotion_task_spec = _task_names.locomotion_task_spec +normalize_locomotion_task_id = _task_names.normalize_locomotion_task_id @dataclass @@ -75,7 +79,7 @@ class BenchRecord: DEFAULT_NJMAX_BY_TASK = { "go1_joystick_flat": 100, "go2_joystick_flat": 100, - "g1_joystick_flat": 150, + "g1_walk_flat": 150, } @@ -108,30 +112,33 @@ def _require_mujoco_warp() -> None: ) -def _load_task_model(task_name: str) -> "mujoco.MjModel": +def _load_task_model(task_name: str) -> Any: cfg = locomotion_task_spec(task_name).config_cls() - return mujoco.MjModel.from_xml_path(cfg.model_file) + return cast(Any, mujoco).MjModel.from_xml_path(cfg.model_file) def _task_njmax(task_name: str) -> int: return DEFAULT_NJMAX_BY_TASK.get(task_name, -1) -def _make_warp_data(model: "mujoco.MjModel", batch_size: int, njmax: int): +def _make_warp_data(model: Any, batch_size: int, njmax: int): + mj_warp_mod = cast(Any, mj_warp) try: if njmax > 0: - return mj_warp.make_data(model, nworld=batch_size, njmax=njmax, nconmax=njmax) - return mj_warp.make_data(model, nworld=batch_size) + return mj_warp_mod.make_data(model, nworld=batch_size, njmax=njmax, nconmax=njmax) + return mj_warp_mod.make_data(model, nworld=batch_size) except TypeError: - return mj_warp.make_data(model, nworld=batch_size) + return mj_warp_mod.make_data(model, nworld=batch_size) def _run_warp(warp_model, warp_data, nstep: int, niter: int) -> float: + mj_warp_mod = cast(Any, mj_warp) + warp_mod = cast(Any, warp) t0 = time.perf_counter() for _ in range(niter): for _ in range(nstep): - mj_warp.step(warp_model, warp_data) - warp.synchronize() + mj_warp_mod.step(warp_model, warp_data) + warp_mod.synchronize() return (time.perf_counter() - t0) / niter @@ -144,7 +151,7 @@ def _bench_one_task( ) -> List[BenchRecord]: task_key = normalize_locomotion_task_id(task_name) model = _load_task_model(task_key) - warp_model = mj_warp.put_model(model) + warp_model = cast(Any, mj_warp).put_model(model) njmax = _task_njmax(task_key) records: List[BenchRecord] = [] diff --git a/benchmark/core/task_names.py b/benchmark/core/task_names.py index fec3dd002..9607cf55c 100644 --- a/benchmark/core/task_names.py +++ b/benchmark/core/task_names.py @@ -2,7 +2,7 @@ from dataclasses import dataclass -from unilab.envs.locomotion.g1.joystick import G1JoystickPPOCfg +from unilab.envs.locomotion.g1.joystick import G1WalkFlatCfg from unilab.envs.locomotion.go1.joystick import Go1JoystickCfg from unilab.envs.locomotion.go2.joystick import Go2JoystickCfg from unilab.envs.manipulation.sharpa_inhand.rotation import SharpaInhandRotationCfg @@ -29,11 +29,11 @@ class LocomotionTaskSpec: display_name="go2_joystick_flat", config_cls=Go2JoystickCfg, ), - "g1_joystick_flat": LocomotionTaskSpec( - owner_task_id="g1_joystick_flat", - env_task_name="G1JoystickFlat", - display_name="g1_joystick_flat", - config_cls=G1JoystickPPOCfg, + "g1_walk_flat": LocomotionTaskSpec( + owner_task_id="g1_walk_flat", + env_task_name="G1WalkFlat", + display_name="g1_walk_flat", + config_cls=G1WalkFlatCfg, ), "sharpa_inhand": LocomotionTaskSpec( owner_task_id="sharpa_inhand", diff --git a/conf/appo/task/g1_joystick_flat/mujoco.yaml b/conf/appo/task/g1_walk_flat/mujoco.yaml similarity index 86% rename from conf/appo/task/g1_joystick_flat/mujoco.yaml rename to conf/appo/task/g1_walk_flat/mujoco.yaml index 0474ba6eb..83fce71ba 100644 --- a/conf/appo/task/g1_joystick_flat/mujoco.yaml +++ b/conf/appo/task/g1_walk_flat/mujoco.yaml @@ -1,10 +1,15 @@ # @package _global_ training: - task_name: G1JoystickFlat + task_name: G1WalkFlat sim_backend: mujoco algo: max_iterations: 500 save_interval: 100 +env: + control_config: + action_scale: 0.25 + curriculum: + enabled: false reward: scales: tracking_lin_vel: 2.0 diff --git a/conf/offpolicy/task/flashsac/g1_walk_flat/motrix.yaml b/conf/offpolicy/task/flashsac/g1_walk_flat/motrix.yaml index ec77aeb95..76903fffc 100644 --- a/conf/offpolicy/task/flashsac/g1_walk_flat/motrix.yaml +++ b/conf/offpolicy/task/flashsac/g1_walk_flat/motrix.yaml @@ -6,6 +6,15 @@ algo: num_envs: 2048 max_iterations: 5000 save_interval: 1000 +env: + curriculum: + enabled: true + initial_scale: 0.5 + min_scale: 0.5 + max_scale: 1.0 + level_down_threshold: 150.0 + level_up_threshold: 750.0 + degree: 0.001 reward: scales: tracking_lin_vel: 2.2 diff --git a/conf/offpolicy/task/flashsac/g1_walk_flat/mujoco.yaml b/conf/offpolicy/task/flashsac/g1_walk_flat/mujoco.yaml index c351cd8a1..009a705e4 100644 --- a/conf/offpolicy/task/flashsac/g1_walk_flat/mujoco.yaml +++ b/conf/offpolicy/task/flashsac/g1_walk_flat/mujoco.yaml @@ -8,13 +8,21 @@ algo: max_iterations: 6000 save_interval: 1000 updates_per_step: 8 - #use_symmetry: true + #use_symmetry: true replay_buffer_n: 1024 env: control_config: action_scale: 1.0 gait_phase_init_mode: "offset_phase" reset_base_qvel_limit: 0.5 + curriculum: + enabled: true + initial_scale: 0.5 + min_scale: 0.5 + max_scale: 1.0 + level_down_threshold: 150.0 + level_up_threshold: 750.0 + degree: 0.001 noise_config: scale_gyro: 0.0 scale_gravity: 0.0 diff --git a/conf/offpolicy/task/flashsac/g1_joystick_flat/mujoco.yaml b/conf/offpolicy/task/flashsac/g1_walk_flat_amp/mujoco.yaml similarity index 94% rename from conf/offpolicy/task/flashsac/g1_joystick_flat/mujoco.yaml rename to conf/offpolicy/task/flashsac/g1_walk_flat_amp/mujoco.yaml index c10eced1d..88d54ab1b 100644 --- a/conf/offpolicy/task/flashsac/g1_joystick_flat/mujoco.yaml +++ b/conf/offpolicy/task/flashsac/g1_walk_flat_amp/mujoco.yaml @@ -1,6 +1,6 @@ # @package _global_ training: - task_name: G1JoystickFlat + task_name: G1WalkFlat sim_backend: mujoco use_amp: true algo: @@ -18,6 +18,8 @@ env: vel_limit: - [-1.0, -0.5, -1.0] - [1.0, 0.5, 1.0] + curriculum: + enabled: false reward: scales: tracking_lin_vel: 1.0 diff --git a/conf/offpolicy/task/sac/g1_walk_flat/motrix.yaml b/conf/offpolicy/task/sac/g1_walk_flat/motrix.yaml index 16b240c04..b28bb3b5e 100644 --- a/conf/offpolicy/task/sac/g1_walk_flat/motrix.yaml +++ b/conf/offpolicy/task/sac/g1_walk_flat/motrix.yaml @@ -15,6 +15,14 @@ env: domain_rand: randomize_kp: false randomize_kd: false + curriculum: + enabled: true + initial_scale: 0.5 + min_scale: 0.5 + max_scale: 1.0 + level_down_threshold: 150.0 + level_up_threshold: 750.0 + degree: 0.001 reward: scales: tracking_lin_vel: 2.2 diff --git a/conf/offpolicy/task/sac/g1_walk_flat/mujoco.yaml b/conf/offpolicy/task/sac/g1_walk_flat/mujoco.yaml index bfb5d556e..bf6901d16 100644 --- a/conf/offpolicy/task/sac/g1_walk_flat/mujoco.yaml +++ b/conf/offpolicy/task/sac/g1_walk_flat/mujoco.yaml @@ -17,6 +17,14 @@ env: action_scale: 1.0 gait_phase_init_mode: "offset_phase" reset_base_qvel_limit: 0.5 + curriculum: + enabled: true + initial_scale: 0.5 + min_scale: 0.5 + max_scale: 1.0 + level_down_threshold: 150.0 + level_up_threshold: 750.0 + degree: 0.001 noise_config: scale_gyro: 0.0 scale_gravity: 0.0 diff --git a/conf/offpolicy/task/sac/g1_walk_rough/mujoco.yaml b/conf/offpolicy/task/sac/g1_walk_rough/mujoco.yaml index 8d07b2c88..fb7938dba 100644 --- a/conf/offpolicy/task/sac/g1_walk_rough/mujoco.yaml +++ b/conf/offpolicy/task/sac/g1_walk_rough/mujoco.yaml @@ -17,6 +17,14 @@ env: action_scale: 1.0 gait_phase_init_mode: "offset_phase" reset_base_qvel_limit: 0.5 + curriculum: + enabled: true + initial_scale: 0.5 + min_scale: 0.5 + max_scale: 1.0 + level_down_threshold: 150.0 + level_up_threshold: 750.0 + degree: 0.001 noise_config: scale_gyro: 0.0 scale_gravity: 0.0 diff --git a/conf/offpolicy/task/td3/g1_walk_flat/mujoco.yaml b/conf/offpolicy/task/td3/g1_walk_flat/mujoco.yaml index 714d0fce9..2705afec6 100644 --- a/conf/offpolicy/task/td3/g1_walk_flat/mujoco.yaml +++ b/conf/offpolicy/task/td3/g1_walk_flat/mujoco.yaml @@ -9,6 +9,14 @@ env: action_scale: 1.0 gait_phase_init_mode: "offset_phase" reset_base_qvel_limit: 0.5 + curriculum: + enabled: true + initial_scale: 0.5 + min_scale: 0.5 + max_scale: 1.0 + level_down_threshold: 150.0 + level_up_threshold: 750.0 + degree: 0.001 noise_config: scale_gyro: 0.0 scale_gravity: 0.0 diff --git a/conf/ppo/task/g1_joystick_flat/motrix.yaml b/conf/ppo/task/g1_walk_flat/motrix.yaml similarity index 95% rename from conf/ppo/task/g1_joystick_flat/motrix.yaml rename to conf/ppo/task/g1_walk_flat/motrix.yaml index 97976b8b9..56e16b725 100644 --- a/conf/ppo/task/g1_joystick_flat/motrix.yaml +++ b/conf/ppo/task/g1_walk_flat/motrix.yaml @@ -1,6 +1,6 @@ # @package _global_ training: - task_name: G1JoystickFlat + task_name: G1WalkFlat sim_backend: motrix algo: num_envs: 2048 @@ -27,6 +27,8 @@ env: - [0.7, 0.0, 0.0] gait_phase_init_mode: offset_phase reset_base_qvel_limit: 0.05 + curriculum: + enabled: false reward: scales: tracking_lin_vel: 2.0 diff --git a/conf/ppo/task/g1_joystick_flat/mujoco.yaml b/conf/ppo/task/g1_walk_flat/mujoco.yaml similarity index 87% rename from conf/ppo/task/g1_joystick_flat/mujoco.yaml rename to conf/ppo/task/g1_walk_flat/mujoco.yaml index e6d68aab4..b47557b30 100644 --- a/conf/ppo/task/g1_joystick_flat/mujoco.yaml +++ b/conf/ppo/task/g1_walk_flat/mujoco.yaml @@ -1,6 +1,6 @@ # @package _global_ training: - task_name: G1JoystickFlat + task_name: G1WalkFlat sim_backend: mujoco algo: num_envs: 2048 @@ -8,6 +8,11 @@ algo: obs_groups: actor: - actor +env: + control_config: + action_scale: 0.25 + curriculum: + enabled: false reward: scales: tracking_lin_vel: 2.0 diff --git a/docs/users/zh_CN/02-simulation-backends.md b/docs/users/zh_CN/02-simulation-backends.md index de3706265..81ba7f2f9 100644 --- a/docs/users/zh_CN/02-simulation-backends.md +++ b/docs/users/zh_CN/02-simulation-backends.md @@ -43,7 +43,7 @@ uv run scripts/generate_support_matrix.py --write |------------|------------|--------|--------| | PPO (torch) | `go1_joystick_flat` (Go1 joystick) | Tested | Tested | | PPO (torch) | `go2_joystick_flat` (Go2 joystick) | Tested | Tested | -| PPO (torch) | `g1_joystick_flat` (G1 joystick) | Tested | Tested | +| PPO (torch) | `g1_walk_flat` (G1 walk flat) | Tested | Tested | | PPO (torch) | `g1_motion_tracking` (G1 motion tracking) | Tested | Tested | | PPO (torch) | `g1_flip_tracking` (G1 flip tracking) | Tested | Tested | | PPO (torch) | `allegro_inhand` (Allegro in-hand) | Tested | Tested | @@ -52,7 +52,7 @@ uv run scripts/generate_support_matrix.py --write | PPO (torch) | `sharpa_inhand_grasp` (sharpa inhand grasp) | Tested | Tested | | PPO (mlx) | `go1_joystick_flat` (Go1 joystick) | Tested | Tested | | PPO (mlx) | `go2_joystick_flat` (Go2 joystick) | Tested | Tested | -| PPO (mlx) | `g1_joystick_flat` (G1 joystick) | Tested | Tested | +| PPO (mlx) | `g1_walk_flat` (G1 walk flat) | Tested | Tested | | PPO (mlx) | `g1_motion_tracking` (G1 motion tracking) | Configured | Configured | | PPO (mlx) | `g1_flip_tracking` (G1 flip tracking) | Configured | Configured | | PPO (mlx) | `allegro_inhand` (Allegro in-hand) | Configured | Configured | @@ -61,7 +61,7 @@ uv run scripts/generate_support_matrix.py --write | PPO (mlx) | `sharpa_inhand_grasp` (sharpa inhand grasp) | Configured | Configured | | APPO (torch) | `go1_joystick_flat` (Go1 joystick) | Tested | Registered | | APPO (torch) | `go2_joystick_flat` (Go2 joystick) | Tested | Registered | -| APPO (torch) | `g1_joystick_flat` (G1 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 | @@ -74,7 +74,7 @@ uv run scripts/generate_support_matrix.py --write - Registry bootstrap: `src/unilab/envs/**` decorators via `unilab.utils.algo_utils.ensure_registries()`. - Owner YAML scan: `conf/ppo/task/**`, `conf/appo/task/**`, `conf/offpolicy/task/**`. - Generic compose coverage: `tests/config/test_config_system.py::test_supported_task_composes`. -- MLX-specific compose coverage only upgrades task owners listed in `tests/config/test_config_system.py::_PPO_MLX_TASKS`: `go1_joystick_flat`, `go2_joystick_flat`, `g1_joystick_flat`. +- MLX-specific compose coverage only upgrades task owners listed in `tests/config/test_config_system.py::_PPO_MLX_TASKS`: `go1_joystick_flat`, `go2_joystick_flat`, `g1_walk_flat`. - MLX runtime smoke: `tests/algos/test_mlx_ppo.py::test_mlx_ppo_one_iteration_real_env` currently exercises `go2_joystick_flat/mujoco`. diff --git a/docs/users/zh_CN/06-domain-randomization.md b/docs/users/zh_CN/06-domain-randomization.md index e6c9d6010..b52bd9d14 100644 --- a/docs/users/zh_CN/06-domain-randomization.md +++ b/docs/users/zh_CN/06-domain-randomization.md @@ -17,7 +17,7 @@ ## 现状结论 1. 当前已接入 DR provider 的任务全部使用统一 DR 入口,没有任务绕开 `DomainRandomizationManager` 直接在 `reset()` 里做另一套 DR 流程。 -2. 形式上基本都是结构化的:任务文件内定义 `domain_rand` 配置 dataclass、`DomainRandomizationProvider`、`ResetPlan`,`G1WalkFlat` 复用 `G1Joystick` 的 provider。 +2. 形式上基本都是结构化的:任务文件内定义 `domain_rand` 配置 dataclass、`DomainRandomizationProvider`、`ResetPlan`,`G1WalkFlat` 复用 `G1Walk` 的 provider。 3. 现在“统一”的主要是入口和执行流程,不是所有随机项本身。公共 helper [`build_common_reset_randomization()`](../../../src/unilab/dr/dr_utils.py) 目前生成 `base_mass_delta`、`base_com_offset`、`kp`、`kd`;公共 interval helper 目前只生成 push。 4. [`ResetRandomizationPayload`](../../../src/unilab/dr/types.py) 已经能表达 `body_iquat`、`body_inertia`、`kp`、`kd`,且 [`MuJoCoBackend`](../../../src/unilab/base/backend/mujoco_backend.py) 已声明支持。是否真正使用这些项,仍取决于任务 provider 是否采样并下发。 5. [`MotrixBackend`](../../../src/unilab/base/backend/motrix_backend.py) 当前支持 `base_mass_delta`、`base_com_offset`、`kp`、`kd` 和 interval push;并在初始化阶段要求模型 actuator 全部为 position actuator。 @@ -29,8 +29,8 @@ | --- | --- | --- | --- | --- | --- | | `Go1JoystickFlat` | 是 | 是:`Domain_Rand + Provider + ResetPlan` | 任务状态采样 + common payload | push | [`go1/joystick.py`](../../../src/unilab/envs/locomotion/go1/joystick.py) | | `Go2JoystickFlat` | 是 | 是:`Domain_Rand + Provider + ResetPlan` | 任务状态采样 + common payload | push | [`go2/joystick.py`](../../../src/unilab/envs/locomotion/go2/joystick.py) | -| `G1JoystickFlat` | 是 | 是:`Domain_Rand + Provider + ResetPlan` | 任务状态采样 + common payload | push | [`g1/joystick.py`](../../../src/unilab/envs/locomotion/g1/joystick.py) | -| `G1WalkFlat` | 是 | 是:复用 [`G1JoystickDomainRandomizationProvider`](../../../src/unilab/envs/locomotion/g1/joystick.py) | 任务状态采样 + common payload | push | [`g1/joystick_sac.py`](../../../src/unilab/envs/locomotion/g1/joystick_sac.py) | +| `G1WalkFlat` | 是 | 是:`Domain_Rand + Provider + ResetPlan` | 任务状态采样 + common payload | push | [`g1/joystick.py`](../../../src/unilab/envs/locomotion/g1/joystick.py) | +| `G1WalkFlat` | 是 | 是:复用 [`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 采样,`randomization=None` | 无 | [`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) | @@ -42,8 +42,8 @@ | --- | --- | --- | --- | | `Go1JoystickFlat` | base xy;base yaw;base qvel;command 采样;`current_actions/last_actions` 清零;可选 `base_mass_delta`;可选 `base_com_offset` | `push_robots` | 默认开启 `base_mass_delta`、`base_com_offset`、push | | `Go2JoystickFlat` | base xy;base yaw;base qvel;command 采样;`current_actions/last_actions` 清零;kp/kd 随机化(默认开启);可选 `base_mass_delta`;可选 `base_com_offset` | `push_robots` | kp/kd 默认开启;common payload 和 push 默认关闭 | -| `G1JoystickFlat` | base xy;base yaw;按 `reset_base_qvel_limit` 采样 base qvel;command 采样;`gait_phase` 采样;`current_actions/last_actions` 清零;kp/kd 随机化(默认开启);可选 `base_mass_delta`;可选 `base_com_offset` | `push_robots` | kp/kd 默认开启;common payload 和 push 默认关闭 | -| `G1WalkFlat` | 与 `G1JoystickFlat` 相同,直接复用同一个 provider | `push_robots` | kp/kd 默认开启;common payload 和 push 默认关闭 | +| `G1WalkFlat` | base xy;base yaw;按 `reset_base_qvel_limit` 采样 base qvel;command 采样;`gait_phase` 采样;`current_actions/last_actions` 清零;kp/kd 随机化(默认开启);可选 `base_mass_delta`;可选 `base_com_offset` | `push_robots` | kp/kd 默认开启;common payload 和 push 默认关闭 | +| `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` | `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 | 无 | grasp cache 路径可用时默认会采样;`joint_noise`、`ball_vel_noise`、`ball_z_offset` 默认 0 | diff --git a/src/unilab/config/locomotion_params.py b/src/unilab/config/locomotion_params.py index a3aa56637..4eb197534 100644 --- a/src/unilab/config/locomotion_params.py +++ b/src/unilab/config/locomotion_params.py @@ -9,7 +9,6 @@ { "Go1JoystickFlat", "Go2JoystickFlat", - "G1JoystickFlat", "G1WalkFlat", "G1WalkRough", "G1MotionTracking", diff --git a/src/unilab/docs/support_matrix.py b/src/unilab/docs/support_matrix.py index ed791b4d7..5fe9931f4 100644 --- a/src/unilab/docs/support_matrix.py +++ b/src/unilab/docs/support_matrix.py @@ -20,22 +20,20 @@ _TASK_ORDER = { "go1_joystick_flat": 0, "go2_joystick_flat": 1, - "g1_joystick_flat": 2, - "g1_motion_tracking": 3, - "g1_flip_tracking": 4, - "g1_walk_flat": 5, - "g1_walk_rough": 6, + "g1_walk_flat": 2, + "g1_walk_rough": 3, + "g1_motion_tracking": 4, + "g1_flip_tracking": 5, "allegro_inhand": 7, "allegro_sac": 8, } _TASK_LABELS = { "go1_joystick_flat": "Go1 joystick", "go2_joystick_flat": "Go2 joystick", - "g1_joystick_flat": "G1 joystick", - "g1_motion_tracking": "G1 motion tracking", - "g1_flip_tracking": "G1 flip tracking", "g1_walk_flat": "G1 walk flat", "g1_walk_rough": "G1 walk rough", + "g1_motion_tracking": "G1 motion tracking", + "g1_flip_tracking": "G1 flip tracking", "allegro_inhand": "Allegro in-hand", "allegro_sac": "Allegro SAC in-hand", } diff --git a/src/unilab/envs/locomotion/g1/__init__.py b/src/unilab/envs/locomotion/g1/__init__.py index fac18b920..b99f8b105 100644 --- a/src/unilab/envs/locomotion/g1/__init__.py +++ b/src/unilab/envs/locomotion/g1/__init__.py @@ -1,7 +1,9 @@ -from .joystick import G1JoystickPPO, G1JoystickPPOCfg -from .joystick_sac import ( - G1WalkFlat, +from .joystick import ( + G1WalkControlConfig, + G1WalkEnv, + G1WalkEnvCfg, G1WalkFlatCfg, - G1WalkRough, + G1WalkLegacyRewardConfig, + G1WalkRewardConfig, G1WalkRoughCfg, ) diff --git a/src/unilab/envs/locomotion/g1/joystick.py b/src/unilab/envs/locomotion/g1/joystick.py index 1023ef34a..57465dbfa 100644 --- a/src/unilab/envs/locomotion/g1/joystick.py +++ b/src/unilab/envs/locomotion/g1/joystick.py @@ -1,4 +1,4 @@ -"""G1 Joystick environments - PPO and SAC variants.""" +"""G1 joystick locomotion environments.""" from __future__ import annotations @@ -13,6 +13,7 @@ from unilab.base import registry from unilab.base.augmentation import SymmetryObsLayout from unilab.base.backend import create_backend +from unilab.base.curriculum import EpisodeLengthTracker, PenaltyCurriculum from unilab.base.dtype_config import get_global_dtype from unilab.base.np_env import NpEnvState from unilab.envs.locomotion.common import rewards @@ -124,7 +125,7 @@ def compute_forward_command_mask(commands: np.ndarray) -> np.ndarray: @dataclass -class RewardConfigPPO: +class G1RewardConfig: scales: dict[str, float] tracking_sigma: float gait_frequency: float @@ -134,6 +135,7 @@ class RewardConfigPPO: min_base_height: float max_tilt_deg: float min_forward_speed_for_gait_reward: float = 0.0 + close_feet_threshold: float = 0.15 pose_weights: list[float] = field( default_factory=lambda: [ 0.01, @@ -169,21 +171,36 @@ class RewardConfigPPO: ) -# PPO Environment -@registry.envcfg("G1JoystickFlat") @dataclass -class G1JoystickPPOCfg(G1BaseCfg): +class G1WalkLegacyRewardConfig(G1RewardConfig): + pass + + +@dataclass +class CurriculumConfig: + enabled: bool = False + initial_scale: float = 0.5 + min_scale: float = 0.5 + max_scale: float = 1.0 + level_down_threshold: float = 150.0 + level_up_threshold: float = 750.0 + degree: float = 0.001 + + +@dataclass +class G1WalkEnvCfg(G1BaseCfg): model_file: str = str(ASSETS_ROOT_PATH / "robots" / "g1" / "scene_flat.xml") max_episode_seconds: float = 20.0 init_state: InitState = field(default_factory=InitState) commands: Commands = field(default_factory=Commands) - reward_config: RewardConfigPPO | None = None + reward_config: G1RewardConfig | None = None domain_rand: G1DomainRandConfig = field(default_factory=G1DomainRandConfig) gait_phase_init_mode: str = "offset_phase" reset_base_qvel_limit: float = 0.5 + curriculum: CurriculumConfig = field(default_factory=CurriculumConfig) -class G1JoystickDomainRandomizationProvider(LocomotionDRProvider): +class G1WalkDomainRandomizationProvider(LocomotionDRProvider): def __init__(self, *, base_kp: np.ndarray | None = None, base_kd: np.ndarray | None = None): self._base_kp = base_kp self._base_kd = base_kd @@ -221,13 +238,11 @@ def _compute_reset_obs( return env._compute_obs(info_updates, linvel, gyro, gravity, dof_pos, dof_vel) # type: ignore[no-any-return] -@registry.env("G1JoystickFlat", sim_backend="mujoco") -@registry.env("G1JoystickFlat", sim_backend="motrix") -class G1JoystickPPO(G1BaseEnv): - _cfg: G1JoystickPPOCfg +class G1WalkEnv(G1BaseEnv): + _cfg: G1WalkEnvCfg _reward_cfg: Any - def __init__(self, cfg: G1JoystickPPOCfg, num_envs=1, backend_type="mujoco"): + def __init__(self, cfg: G1WalkEnvCfg, num_envs=1, backend_type="mujoco"): if cfg.reward_config is None: raise ValueError("reward_config must be provided via Hydra configuration") backend = create_backend( @@ -249,13 +264,27 @@ def __init__(self, cfg: G1JoystickPPOCfg, num_envs=1, backend_type="mujoco"): if self._pose_weights.shape[0] != self._num_action: raise ValueError("pose_weights length mismatch") self._upper_body_pose_weights = build_upper_body_pose_weights(self._reward_cfg.pose_weights) + self._episode_tracker: EpisodeLengthTracker | None = None + self._penalty_curriculum: PenaltyCurriculum | None = None + if cfg.curriculum.enabled: + self._episode_tracker = EpisodeLengthTracker(num_envs) + self._penalty_curriculum = PenaltyCurriculum( + self, + enabled=True, + initial_scale=cfg.curriculum.initial_scale, + min_scale=cfg.curriculum.min_scale, + max_scale=cfg.curriculum.max_scale, + level_down_threshold=cfg.curriculum.level_down_threshold, + level_up_threshold=cfg.curriculum.level_up_threshold, + degree=cfg.curriculum.degree, + ) self._init_reward_functions() if cfg.domain_rand.randomize_kp or cfg.domain_rand.randomize_kd: base_kp, base_kd = backend.get_actuator_gains() - dr_provider = G1JoystickDomainRandomizationProvider(base_kp=base_kp, base_kd=base_kd) + dr_provider = G1WalkDomainRandomizationProvider(base_kp=base_kp, base_kd=base_kd) else: - dr_provider = G1JoystickDomainRandomizationProvider() + dr_provider = G1WalkDomainRandomizationProvider() self._init_domain_randomization(dr_provider) @property @@ -271,16 +300,22 @@ def _init_reward_functions(self): "under_speed": rewards.under_speed, "lin_vel_z": rewards.lin_vel_z, "orientation": rewards.orientation, + "penalty_orientation": rewards.orientation, "ang_vel_xy": rewards.ang_vel_xy, + "penalty_ang_vel_xy": rewards.ang_vel_xy, "action_rate": rewards.action_rate, + "penalty_action_rate": rewards.action_rate, "base_height": rewards.base_height, "pose": rewards.weighted_pose, "upper_body_pose": self._reward_upper_body_pose, + "penalty_close_feet_xy": self._reward_close_feet_xy, "penalty_feet_ori": self._reward_feet_ori, "feet_phase": self._reward_feet_phase, "feet_phase_contrast": self._reward_feet_phase_contrast, "feet_phase_contact": self._reward_feet_phase_contact, "feet_double_stance": self._reward_feet_double_stance, + "feet_air_time": self._reward_feet_air_time, + "alive": rewards.alive, } def update_state(self, state: NpEnvState) -> NpEnvState: @@ -299,29 +334,111 @@ def update_state(self, state: NpEnvState) -> NpEnvState: reward = self._compute_reward(state.info, linvel, gyro, gravity, dof_pos, dof_vel) obs = self._compute_obs(state.info, linvel, gyro, gravity, dof_pos, dof_vel) - return state.replace(obs=obs, reward=reward, terminated=terminated) + state = state.replace(obs=obs, reward=reward, terminated=terminated) + + if ( + self._episode_tracker is None + or self._penalty_curriculum is None + or not np.any(state.done) + ): + return state + + done_indices = np.where(state.done)[0] + episode_lengths = state.info["steps"][done_indices] + 1 + self._episode_tracker.update(episode_lengths) + self._penalty_curriculum.update(self._episode_tracker.average_length) + + if "log" not in state.info: + state.info["log"] = {} + state.info["log"]["curriculum/average_episode_length"] = float( + self._episode_tracker.average_length + ) + state.info["log"]["curriculum/penalty_scale"] = float( + self._penalty_curriculum.current_scale + ) + return state def _compute_obs( self, info: dict, linvel, gyro, gravity, dof_pos, dof_vel ) -> dict[str, np.ndarray]: noise_cfg = self._cfg.noise_config diff = dof_pos - self.default_angles - gyro = self._obs_noise(gyro, noise_cfg.scale_gyro) - gravity = self._obs_noise(gravity, noise_cfg.scale_gravity) - diff = self._obs_noise(diff, noise_cfg.scale_joint_angle) - dof_vel = self._obs_noise(dof_vel, noise_cfg.scale_joint_vel) - linvel = self._obs_noise(linvel, noise_cfg.scale_linvel) command = info["commands"] last_actions = info.get("current_actions", np.zeros_like(diff)) gait_phase = info.get("gait_phase", np.zeros((self._num_envs, 2), dtype=get_global_dtype())) + walk_profile = self._uses_walk_observation_profile() + + noisy_gyro = self._obs_noise(gyro, noise_cfg.scale_gyro) + noisy_gravity = self._obs_noise(gravity, noise_cfg.scale_gravity) + noisy_diff = self._obs_noise(diff, noise_cfg.scale_joint_angle) + noisy_dof_vel = self._obs_noise(dof_vel, noise_cfg.scale_joint_vel) + actor_gyro_scale = 0.25 if walk_profile else 1.0 + actor_dof_vel_scale = 0.05 if walk_profile else 1.0 + actor = np.concatenate( - [gyro, -gravity, diff, dof_vel, last_actions, command, gait_phase], + [ + noisy_gyro * actor_gyro_scale, + -noisy_gravity, + noisy_diff, + noisy_dof_vel * actor_dof_vel_scale, + last_actions, + command, + gait_phase, + ], + axis=1, + dtype=get_global_dtype(), + ) + + critic_gyro_scale = 0.25 if walk_profile else 1.0 + critic_dof_vel_scale = 0.05 if walk_profile else 1.0 + critic_linvel_scale = 2.0 if walk_profile else 1.0 + critic_base = np.concatenate( + [ + gyro * critic_gyro_scale, + -gravity, + diff, + dof_vel * critic_dof_vel_scale, + last_actions, + command, + gait_phase, + ], + axis=1, + dtype=get_global_dtype(), + ) + critic = np.concatenate( + [ + critic_base, + np.asarray(linvel * critic_linvel_scale, dtype=get_global_dtype()), + ], axis=1, dtype=get_global_dtype(), ) - critic = np.concatenate([actor, linvel], axis=1, dtype=get_global_dtype()) + return {"obs": actor, "critic": critic} + def _uses_walk_observation_profile(self) -> bool: + scales = getattr(getattr(self, "_reward_cfg", None), "scales", None) + if scales is None: + reward_cfg = getattr(self._cfg, "reward_config", None) + scales = getattr(reward_cfg, "scales", None) + + if scales is not None: + if any( + key in scales + for key in ( + "penalty_orientation", + "penalty_ang_vel_xy", + "penalty_action_rate", + "alive", + ) + ): + return True + if any(key in scales for key in ("orientation", "ang_vel_xy", "action_rate")): + return False + + curriculum = getattr(self._cfg, "curriculum", None) + return bool(curriculum is not None and curriculum.enabled) + def _actor_symmetry_obs_layout(self) -> SymmetryObsLayout: return ( ("gyro", 3), @@ -334,7 +451,11 @@ def _actor_symmetry_obs_layout(self) -> SymmetryObsLayout: ) def get_symmetry_obs_layouts(self) -> dict[str, SymmetryObsLayout]: - return {"obs": self._actor_symmetry_obs_layout()} + actor_layout = self._actor_symmetry_obs_layout() + return { + "obs": actor_layout, + "critic": (*actor_layout, ("linvel", 3)), + } def build_symmetry_augmentation(self, *, device: str): if self._backend.backend_type != "mujoco": @@ -455,6 +576,23 @@ def _reward_feet_ori(self, ctx: RewardContext): + np.square(right_foot_quat[:, 2]) ) + def _reward_close_feet_xy(self, ctx: RewardContext): + left_foot = self._backend.get_sensor_data("left_foot_pos") + right_foot = self._backend.get_sensor_data("right_foot_pos") + feet_dist = np.linalg.norm(left_foot[:, :2] - right_foot[:, :2], axis=1) + return np.where( + feet_dist < self._reward_cfg.close_feet_threshold, + np.square(feet_dist - self._reward_cfg.close_feet_threshold), + 0.0, + ) + + def _reward_feet_air_time(self, ctx: RewardContext): + air_time = ctx.info.get( + "feet_air_time", np.zeros((self._num_envs, 2), dtype=get_global_dtype()) + ) + in_range = (air_time > 0.05) & (air_time < 0.5) + return np.sum(in_range.astype(float), axis=1) + def _reward_upper_body_pose(self, ctx: RewardContext): diff = ctx.dof_pos - self.default_angles return np.asarray( @@ -475,3 +613,47 @@ def apply_action(self, actions: np.ndarray, state: NpEnvState) -> np.ndarray: ctrl: np.ndarray = actions * self._cfg.control_config.action_scale + self.default_angles return ctrl + + +def _walk_curriculum() -> CurriculumConfig: + return CurriculumConfig( + enabled=True, + initial_scale=0.5, + min_scale=0.5, + max_scale=1.0, + level_down_threshold=150.0, + level_up_threshold=750.0, + degree=0.001, + ) + + +@dataclass +class G1WalkControlConfig: + action_scale: float = 1.0 + simulate_action_latency: bool = False + + +@dataclass +class G1WalkRewardConfig(G1RewardConfig): + """对齐 holosoma G1 walking 奖励权重。""" + + +@registry.envcfg("G1WalkFlat") +@dataclass +class G1WalkFlatCfg(G1WalkEnvCfg): + reward_config: G1WalkRewardConfig | None = None + model_file: str = str(ASSETS_ROOT_PATH / "robots" / "g1" / "scene_flat.xml") + control_config: G1WalkControlConfig = field(default_factory=G1WalkControlConfig) # type: ignore[assignment] + curriculum: CurriculumConfig = field(default_factory=_walk_curriculum) + + +@registry.envcfg("G1WalkRough") +@dataclass +class G1WalkRoughCfg(G1WalkFlatCfg): + model_file: str = str(ASSETS_ROOT_PATH / "robots" / "g1" / "scene_rough.xml") + + +registry.register_env("G1WalkFlat", G1WalkEnv, sim_backend="mujoco") +registry.register_env("G1WalkFlat", G1WalkEnv, sim_backend="motrix") +registry.register_env("G1WalkRough", G1WalkEnv, sim_backend="mujoco") +registry.register_env("G1WalkRough", G1WalkEnv, sim_backend="motrix") diff --git a/src/unilab/envs/locomotion/g1/joystick_sac.py b/src/unilab/envs/locomotion/g1/joystick_sac.py deleted file mode 100644 index 939ceb391..000000000 --- a/src/unilab/envs/locomotion/g1/joystick_sac.py +++ /dev/null @@ -1,242 +0,0 @@ -"""G1 SAC environment - inherits from PPO for code reuse.""" - -from __future__ import annotations - -from dataclasses import dataclass, field - -import numpy as np -from etils import epath - -from unilab.assets import ASSETS_ROOT_PATH -from unilab.base import registry -from unilab.base.augmentation import SymmetryObsLayout -from unilab.base.backend import create_backend -from unilab.base.curriculum import EpisodeLengthTracker, PenaltyCurriculum -from unilab.base.dtype_config import get_global_dtype -from unilab.envs.locomotion.common import rewards -from unilab.envs.locomotion.common.commands import Commands -from unilab.envs.locomotion.common.rewards import RewardContext -from unilab.envs.locomotion.g1.base import G1BaseCfg, G1BaseEnv -from unilab.envs.locomotion.g1.joystick import ( - G1DomainRandConfig, - G1JoystickDomainRandomizationProvider, - G1JoystickPPO, - InitState, -) - - -@dataclass -class ControlConfigSAC: - action_scale: float = 1.0 - simulate_action_latency: bool = False - - -@dataclass -class RewardConfigSAC: - """对齐 holosoma G1 FastSAC 奖励权重""" - - scales: dict[str, float] - tracking_sigma: float - base_height_target: float - min_base_height: float - max_tilt_deg: float - gait_frequency: float - feet_phase_swing_height: float - feet_phase_tracking_sigma: float - close_feet_threshold: float - pose_weights: list[float] - - -@registry.envcfg("G1WalkFlat") -@dataclass -class G1WalkFlatCfg(G1BaseCfg): - reward_config: RewardConfigSAC | None = None - model_file: str = str(ASSETS_ROOT_PATH / "robots" / "g1" / "scene_flat.xml") - max_episode_seconds: float = 20.0 - init_state: InitState = field(default_factory=InitState) - commands: Commands = field(default_factory=Commands) - control_config: ControlConfigSAC = field(default_factory=ControlConfigSAC) # type: ignore[assignment] - domain_rand: G1DomainRandConfig = field(default_factory=G1DomainRandConfig) - gait_phase_init_mode: str = "offset_phase" - reset_base_qvel_limit: float = 0.5 - - -@registry.envcfg("G1WalkRough") -@dataclass -class G1WalkRoughCfg(G1WalkFlatCfg): - model_file: str = str(ASSETS_ROOT_PATH / "robots" / "g1" / "scene_rough.xml") - - -@registry.env("G1WalkFlat", sim_backend="mujoco") -@registry.env("G1WalkFlat", sim_backend="motrix") -class G1WalkFlat(G1JoystickPPO): - """G1 SAC environment - inherits from PPO, overrides rewards.""" - - def __init__(self, cfg: G1WalkFlatCfg, num_envs=1, backend_type="mujoco"): - if cfg.reward_config is None: - raise ValueError("reward_config must be provided via Hydra configuration") - backend = create_backend( - backend_type, - cfg.model_file, - num_envs, - cfg.sim_dt, - base_name=cfg.asset.base_name, - iterations=cfg.iterations, - ) - G1BaseEnv.__init__(self, cfg, backend, num_envs) - self._enable_reward_log = True - self._reward_cfg = cfg.reward_config - self._gait_phase_delta = float(2.0 * np.pi * cfg.reward_config.gait_frequency * cfg.ctrl_dt) - self._pose_weights = np.array(cfg.reward_config.pose_weights, dtype=get_global_dtype()) - - # Curriculum learning - 更宽松的初始配置 - self._episode_tracker = EpisodeLengthTracker(num_envs) - self._penalty_curriculum = PenaltyCurriculum( - self, - enabled=True, - initial_scale=0.5, - min_scale=0.5, - max_scale=1.0, - level_down_threshold=150.0, - level_up_threshold=750.0, - degree=0.001, - ) - - self._init_reward_functions() - if cfg.domain_rand.randomize_kp or cfg.domain_rand.randomize_kd: - base_kp, base_kd = backend.get_actuator_gains() - dr_provider = G1JoystickDomainRandomizationProvider(base_kp=base_kp, base_kd=base_kd) - else: - dr_provider = G1JoystickDomainRandomizationProvider() - self._init_domain_randomization(dr_provider) - - @property - def obs_groups_spec(self) -> dict[str, int]: - # actor trunk = 98, critic path = clean 98-dim trunk + linvel (3). - return {"obs": 98, "critic": 101} - - def get_symmetry_obs_layouts(self) -> dict[str, SymmetryObsLayout]: - actor_layout = self._actor_symmetry_obs_layout() - return { - "obs": actor_layout, - "critic": (*actor_layout, ("linvel", 3)), - } - - def _compute_obs( - self, info: dict, linvel, gyro, gravity, dof_pos, dof_vel - ) -> dict[str, np.ndarray]: - noise_cfg = self._cfg.noise_config - diff = dof_pos - self.default_angles - command = info["commands"] - last_actions = info.get("current_actions", np.zeros_like(diff)) - gait_phase = info.get("gait_phase", np.zeros((self._num_envs, 2), dtype=get_global_dtype())) - - # Clean critic trunk: same layout as actor but noise-free - clean_critic = np.concatenate( - [gyro * 0.25, -gravity, diff, dof_vel * 0.05, last_actions, command, gait_phase], - axis=1, - dtype=get_global_dtype(), - ) - - noisy_gyro = self._obs_noise(gyro, noise_cfg.scale_gyro) - noisy_gravity = self._obs_noise(gravity, noise_cfg.scale_gravity) - noisy_diff = self._obs_noise(diff, noise_cfg.scale_joint_angle) - noisy_dof_vel = self._obs_noise(dof_vel, noise_cfg.scale_joint_vel) - actor = np.concatenate( - [ - noisy_gyro * 0.25, - -noisy_gravity, - noisy_diff, - noisy_dof_vel * 0.05, - last_actions, - command, - gait_phase, - ], - axis=1, - dtype=get_global_dtype(), - ) - - return { - "obs": actor, - "critic": np.concatenate( - [clean_critic, np.asarray(linvel * 2.0, dtype=get_global_dtype())], - axis=1, - dtype=get_global_dtype(), - ), - } - - def _init_reward_functions(self): - """对齐 holosoma G1 FastSAC 奖励函数""" - from typing import Any - - self._reward_fns: dict[str, Any] = { - "tracking_lin_vel": rewards.tracking_lin_vel, - "tracking_ang_vel": rewards.tracking_ang_vel, - "penalty_ang_vel_xy": rewards.ang_vel_xy, - "penalty_orientation": rewards.orientation, - "penalty_action_rate": rewards.action_rate, - "pose": rewards.weighted_pose, - # "penalty_close_feet_xy": self._reward_close_feet_xy, - "penalty_feet_ori": self._reward_feet_ori, - "feet_phase": self._reward_feet_phase, - "alive": rewards.alive, - } - - def _reward_close_feet_xy(self, ctx: RewardContext): - """惩罚双脚过近""" - left_foot = self._backend.get_sensor_data("left_foot_pos") - right_foot = self._backend.get_sensor_data("right_foot_pos") - feet_dist = np.linalg.norm(left_foot[:, :2] - right_foot[:, :2], axis=1) - threshold = self._cfg.reward_config.close_feet_threshold # type: ignore[union-attr] - return np.where(feet_dist < threshold, np.square(feet_dist - threshold), 0.0) - - def _reward_feet_ori(self, ctx: RewardContext): - """惩罚脚部姿态偏差""" - left_foot_quat = self._backend.get_sensor_data("left_foot_quat") - right_foot_quat = self._backend.get_sensor_data("right_foot_quat") - return ( - np.square(left_foot_quat[:, 1]) - + np.square(left_foot_quat[:, 2]) - + np.square(right_foot_quat[:, 1]) - + np.square(right_foot_quat[:, 2]) - ) - - def _reward_feet_air_time(self, ctx: RewardContext): - """奖励脚离地时间""" - air_time = ctx.info.get( - "feet_air_time", np.zeros((self._num_envs, 2), dtype=get_global_dtype()) - ) - in_range = (air_time > 0.05) & (air_time < 0.5) - return np.sum(in_range.astype(float), axis=1) - - def update_state(self, state): - """Override to add curriculum update.""" - # Call parent first to compute terminated/truncated - state = super().update_state(state) - - # Track episode lengths AFTER parent update (when terminated is set) - # Note: steps will be incremented in np_env.step() after this returns - if np.any(state.done): - done_indices = np.where(state.done)[0] - # Add 1 because steps will be incremented after update_state - episode_lengths = state.info["steps"][done_indices] + 1 - self._episode_tracker.update(episode_lengths) - self._penalty_curriculum.update(self._episode_tracker.average_length) - - # Always log curriculum metrics when episode ends - if "log" not in state.info: - state.info["log"] = {} - state.info["log"]["curriculum/average_episode_length"] = float( - self._episode_tracker.average_length - ) - state.info["log"]["curriculum/penalty_scale"] = float( - self._penalty_curriculum.current_scale - ) - - return state - - -@registry.env("G1WalkRough", sim_backend="mujoco") -@registry.env("G1WalkRough", sim_backend="motrix") -class G1WalkRough(G1WalkFlat): - pass diff --git a/tests/algos/test_rsl_rl_runner.py b/tests/algos/test_rsl_rl_runner.py index 054c3ecf0..b8fcafdb9 100644 --- a/tests/algos/test_rsl_rl_runner.py +++ b/tests/algos/test_rsl_rl_runner.py @@ -110,7 +110,7 @@ def get_privileged_observations(self): "env_name", [ "Go2JoystickFlat", - "G1JoystickFlat", + "G1WalkFlat", "AllegroInhandRotation", ], ) diff --git a/tests/base/test_reward_override.py b/tests/base/test_reward_override.py index d673482d3..3ad5a7342 100644 --- a/tests/base/test_reward_override.py +++ b/tests/base/test_reward_override.py @@ -1,5 +1,7 @@ """Test reward config override through registry.""" +from typing import Any, cast + import pytest from unilab.base import registry @@ -18,11 +20,14 @@ def test_reward_override_go1(): base_height_target=0.5, ) - env = registry.make( - "Go1JoystickFlat", - num_envs=1, - sim_backend="mujoco", - env_cfg_override={"reward_config": override_config}, + env = cast( + Any, + registry.make( + "Go1JoystickFlat", + num_envs=1, + sim_backend="mujoco", + env_cfg_override={"reward_config": override_config}, + ), ) assert env._cfg.reward_config.scales["tracking_lin_vel"] == 999.0 @@ -33,9 +38,9 @@ def test_reward_override_g1(): """Test G1 reward config override.""" ensure_registries() - from unilab.envs.locomotion.g1.joystick_sac import RewardConfigSAC + from unilab.envs.locomotion.g1.joystick import G1WalkRewardConfig - override_config = RewardConfigSAC( + override_config = G1WalkRewardConfig( scales={"tracking_lin_vel": 888.0, "alive": 20.0}, tracking_sigma=0.3, base_height_target=0.8, @@ -48,11 +53,14 @@ def test_reward_override_g1(): pose_weights=[0.01] * 29, ) - env = registry.make( - "G1WalkFlat", - num_envs=1, - sim_backend="mujoco", - env_cfg_override={"reward_config": override_config}, + env = cast( + Any, + registry.make( + "G1WalkFlat", + num_envs=1, + sim_backend="mujoco", + env_cfg_override={"reward_config": override_config}, + ), ) assert env._cfg.reward_config.scales["tracking_lin_vel"] == 888.0 diff --git a/tests/config/test_config_system.py b/tests/config/test_config_system.py index 3bd5dfb7f..b53f3f737 100644 --- a/tests/config/test_config_system.py +++ b/tests/config/test_config_system.py @@ -17,7 +17,7 @@ from omegaconf import OmegaConf CONF_DIR = Path(__file__).parent.parent.parent / "conf" -_PPO_MLX_TASKS = {"go1_joystick_flat", "go2_joystick_flat", "g1_joystick_flat"} +_PPO_MLX_TASKS = {"go1_joystick_flat", "go2_joystick_flat", "g1_walk_flat"} def _compose(algo_dir: str, config_name: str = "config", overrides: list[str] | None = None): @@ -205,8 +205,8 @@ def test_offpolicy_g1_walk_flat_motrix_preserves_backend_specific_algo_value(): def test_ppo_g1_backend_specific_hyperparams_remain_separate(): - mujoco_cfg = _compose("ppo", overrides=["task=g1_joystick_flat/mujoco"]) - motrix_cfg = _compose("ppo", overrides=["task=g1_joystick_flat/motrix"]) + mujoco_cfg = _compose("ppo", overrides=["task=g1_walk_flat/mujoco"]) + motrix_cfg = _compose("ppo", overrides=["task=g1_walk_flat/motrix"]) assert mujoco_cfg.algo.max_iterations == 220 assert mujoco_cfg.algo.empirical_normalization is False @@ -269,7 +269,7 @@ def test_offpolicy_g1_walk_flat_motrix_preserves_backend_env_overrides(): def test_cli_override_beats_task_defaults(): cfg = _compose( "ppo", - overrides=["task=g1_joystick_flat/motrix", "algo.max_iterations=1"], + overrides=["task=g1_walk_flat/motrix", "algo.max_iterations=1"], ) assert cfg.algo.max_iterations == 1 diff --git a/tests/config/test_locomotion_params.py b/tests/config/test_locomotion_params.py index 51c0ade1d..c27505c85 100644 --- a/tests/config/test_locomotion_params.py +++ b/tests/config/test_locomotion_params.py @@ -3,6 +3,7 @@ from __future__ import annotations from pathlib import Path +from typing import Any, cast import pytest @@ -165,7 +166,7 @@ def test_offpolicy_g1_rough_terrain_task_overrides(): from hydra import compose, initialize_config_dir from hydra.core.global_hydra import GlobalHydra - from unilab.envs.locomotion.g1.joystick_sac import G1WalkRoughCfg + from unilab.envs.locomotion.g1.joystick import G1WalkRoughCfg GlobalHydra.instance().clear() with initialize_config_dir(config_dir=str(CONF_DIR / "offpolicy"), version_base="1.3"): @@ -179,7 +180,7 @@ def test_offpolicy_g1_rough_terrain_task_overrides(): assert G1WalkRoughCfg().model_file.endswith("scene_rough.xml") -def test_offpolicy_flashsac_g1_joystick_flat_task_overrides(): +def test_offpolicy_flashsac_g1_walk_flat_amp_task_overrides(): from hydra import compose, initialize_config_dir from hydra.core.global_hydra import GlobalHydra @@ -187,10 +188,10 @@ def test_offpolicy_flashsac_g1_joystick_flat_task_overrides(): with initialize_config_dir(config_dir=str(CONF_DIR / "offpolicy"), version_base="1.3"): cfg = compose( "config", - overrides=["algo=flashsac", "task=flashsac/g1_joystick_flat/mujoco"], + overrides=["algo=flashsac", "task=flashsac/g1_walk_flat_amp/mujoco"], ) assert cfg.algo.algo == "flashsac" - assert cfg.training.task_name == "G1JoystickFlat" + assert cfg.training.task_name == "G1WalkFlat" assert cfg.training.sim_backend == "mujoco" assert cfg.training.use_amp is True assert cfg.algo.num_envs == 1024 @@ -202,10 +203,47 @@ def test_offpolicy_flashsac_g1_joystick_flat_task_overrides(): assert cfg.algo.gamma == pytest.approx(0.97) assert cfg.env.control_config.action_scale == pytest.approx(0.5) assert cfg.env.commands.vel_limit[0] == [-1.0, -0.5, -1.0] + assert "obs_profile" not in cfg.env + assert cfg.env.curriculum.enabled is False assert cfg.reward.scales.tracking_ang_vel == pytest.approx(0.75) assert cfg.reward.scales.base_height == pytest.approx(0.0) +def test_g1_task_owner_yamls_preserve_legacy_and_walk_observation_profiles(): + from hydra import compose, initialize_config_dir + from hydra.core.global_hydra import GlobalHydra + + from unilab.envs.locomotion.g1.joystick import G1WalkEnv + + def uses_walk_profile(config_group: str, overrides: list[str]) -> bool: + GlobalHydra.instance().clear() + with initialize_config_dir(config_dir=str(CONF_DIR / config_group), version_base="1.3"): + cfg = compose("config", overrides=overrides) + env = cast(Any, object.__new__(G1WalkEnv)) + env._cfg = cfg.env + env._reward_cfg = cfg.reward + return bool(env._uses_walk_observation_profile()) + + assert uses_walk_profile("ppo", ["task=g1_walk_flat/mujoco"]) is False + assert uses_walk_profile("appo", ["task=g1_walk_flat/mujoco"]) is False + assert ( + uses_walk_profile( + "offpolicy", + ["algo=flashsac", "task=flashsac/g1_walk_flat_amp/mujoco"], + ) + is False + ) + + assert uses_walk_profile("offpolicy", ["algo=sac", "task=sac/g1_walk_flat/mujoco"]) is True + assert uses_walk_profile("offpolicy", ["algo=sac", "task=sac/g1_walk_flat/motrix"]) is True + assert uses_walk_profile("offpolicy", ["algo=sac", "task=sac/g1_walk_rough/mujoco"]) is True + assert uses_walk_profile("offpolicy", ["algo=td3", "task=td3/g1_walk_flat/mujoco"]) is True + assert ( + uses_walk_profile("offpolicy", ["algo=flashsac", "task=flashsac/g1_walk_flat/mujoco"]) + is True + ) + + # --------------------------------------------------------------------------- # Hydra YAML loading — appo # --------------------------------------------------------------------------- @@ -228,10 +266,12 @@ def test_appo_g1_task_overrides(): GlobalHydra.instance().clear() with initialize_config_dir(config_dir=str(CONF_DIR / "appo"), version_base="1.3"): - cfg = compose("config", overrides=["task=g1_joystick_flat/mujoco"]) + cfg = compose("config", overrides=["task=g1_walk_flat/mujoco"]) assert cfg.algo.max_iterations == 500 assert cfg.algo.save_interval == 100 - assert cfg.training.task_name == "G1JoystickFlat" + assert cfg.training.task_name == "G1WalkFlat" + assert "obs_profile" not in cfg.env + assert cfg.env.curriculum.enabled is False # --------------------------------------------------------------------------- @@ -256,9 +296,12 @@ def test_ppo_g1_num_envs(): GlobalHydra.instance().clear() with initialize_config_dir(config_dir=str(CONF_DIR / "ppo"), version_base="1.3"): - cfg = compose("config", overrides=["task=g1_joystick_flat/mujoco"]) + cfg = compose("config", overrides=["task=g1_walk_flat/mujoco"]) assert cfg.algo.num_envs == 2048 assert cfg.algo.max_iterations == 220 + assert cfg.training.task_name == "G1WalkFlat" + assert "obs_profile" not in cfg.env + assert cfg.env.curriculum.enabled is False def test_ppo_go2_num_envs(): diff --git a/tests/config/test_reward_injection.py b/tests/config/test_reward_injection.py index 18cb86252..0135528c6 100644 --- a/tests/config/test_reward_injection.py +++ b/tests/config/test_reward_injection.py @@ -49,7 +49,7 @@ def test_reward_config_conversion(): ensure_registries() - # Test G1 SAC config - registry auto-converts dict to RewardConfigSAC + # Test G1 walk config - registry auto-converts dict to G1WalkRewardConfig g1_dict = { "scales": {"tracking_lin_vel": 2.0, "alive": 10.0}, "tracking_sigma": 0.25, diff --git a/tests/envs/locomotion/g1/test_issue175_regression.py b/tests/envs/locomotion/g1/test_issue175_regression.py new file mode 100644 index 000000000..29a8d7324 --- /dev/null +++ b/tests/envs/locomotion/g1/test_issue175_regression.py @@ -0,0 +1,265 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Any, cast + +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 import registry +from unilab.training.backend_adapter import BackendAdapter +from unilab.utils.algo_utils import ensure_registries + +ROOT_DIR = Path(__file__).parents[4] +CONF_DIR = ROOT_DIR / "conf" + + +_G1_OWNER_CASES = [ + { + "id": "ppo_mujoco", + "config_group": "ppo", + "overrides": ["task=g1_walk_flat/mujoco"], + "task_name": "G1WalkFlat", + "backend": "mujoco", + "profile": "legacy", + "action_scale": 0.25, + "curriculum_enabled": False, + }, + { + "id": "ppo_motrix", + "config_group": "ppo", + "overrides": ["task=g1_walk_flat/motrix"], + "task_name": "G1WalkFlat", + "backend": "motrix", + "profile": "legacy", + "action_scale": 0.5, + "curriculum_enabled": False, + }, + { + "id": "appo_mujoco", + "config_group": "appo", + "overrides": ["task=g1_walk_flat/mujoco"], + "task_name": "G1WalkFlat", + "backend": "mujoco", + "profile": "legacy", + "action_scale": 0.25, + "curriculum_enabled": False, + }, + { + "id": "sac_mujoco", + "config_group": "offpolicy", + "overrides": ["algo=sac", "task=sac/g1_walk_flat/mujoco"], + "task_name": "G1WalkFlat", + "backend": "mujoco", + "profile": "walk", + "action_scale": 1.0, + "curriculum_enabled": True, + }, + { + "id": "sac_motrix", + "config_group": "offpolicy", + "overrides": ["algo=sac", "task=sac/g1_walk_flat/motrix"], + "task_name": "G1WalkFlat", + "backend": "motrix", + "profile": "walk", + "action_scale": 1.0, + "curriculum_enabled": True, + }, + { + "id": "sac_rough", + "config_group": "offpolicy", + "overrides": ["algo=sac", "task=sac/g1_walk_rough/mujoco"], + "task_name": "G1WalkRough", + "backend": "mujoco", + "profile": "walk", + "action_scale": 1.0, + "curriculum_enabled": True, + "model_suffix": "scene_rough.xml", + }, + { + "id": "td3_mujoco", + "config_group": "offpolicy", + "overrides": ["algo=td3", "task=td3/g1_walk_flat/mujoco"], + "task_name": "G1WalkFlat", + "backend": "mujoco", + "profile": "walk", + "action_scale": 1.0, + "curriculum_enabled": True, + }, + { + "id": "flashsac_walk_mujoco", + "config_group": "offpolicy", + "overrides": ["algo=flashsac", "task=flashsac/g1_walk_flat/mujoco"], + "task_name": "G1WalkFlat", + "backend": "mujoco", + "profile": "walk", + "action_scale": 1.0, + "curriculum_enabled": True, + }, + { + "id": "flashsac_walk_motrix", + "config_group": "offpolicy", + "overrides": ["algo=flashsac", "task=flashsac/g1_walk_flat/motrix"], + "task_name": "G1WalkFlat", + "backend": "motrix", + "profile": "walk", + "action_scale": 1.0, + "curriculum_enabled": True, + }, + { + "id": "flashsac_amp_mujoco", + "config_group": "offpolicy", + "overrides": ["algo=flashsac", "task=flashsac/g1_walk_flat_amp/mujoco"], + "task_name": "G1WalkFlat", + "backend": "mujoco", + "profile": "legacy", + "action_scale": 0.5, + "curriculum_enabled": False, + }, +] + + +def _compose_cfg(config_group: str, overrides: list[str]): + GlobalHydra.instance().clear() + with initialize_config_dir(config_dir=str(CONF_DIR / config_group), version_base="1.3"): + return compose("config", overrides=overrides) + + +def _materialize_env_cfg(cfg: Any): + from unilab.envs.locomotion.g1.joystick import G1WalkFlatCfg, G1WalkRoughCfg + + env_cfg_cls = G1WalkRoughCfg if cfg.training.task_name == "G1WalkRough" else G1WalkFlatCfg + return OmegaConf.merge(OmegaConf.structured(env_cfg_cls()), cfg.env) + + +def _build_probe_env(cfg: Any): + from unilab.envs.locomotion.g1.joystick import G1WalkEnv + + env = cast(Any, object.__new__(G1WalkEnv)) + env._num_envs = 1 + env._cfg = _materialize_env_cfg(cfg) + env._reward_cfg = cfg.reward + env.default_angles = np.zeros((1, 29), dtype=np.float32) + env._obs_noise = lambda data, scale: np.asarray(data + 100.0, dtype=np.float32) + return env + + +def _compute_probe_obs(cfg: Any) -> dict[str, np.ndarray]: + env = _build_probe_env(cfg) + return cast( + dict[str, np.ndarray], + env._compute_obs( + { + "commands": np.array([[0.7, 0.0, 0.2]], dtype=np.float32), + "current_actions": np.zeros((1, 29), dtype=np.float32), + "gait_phase": np.array([[0.3, 3.4]], dtype=np.float32), + }, + linvel=np.array([[1.0, 2.0, 3.0]], dtype=np.float32), + gyro=np.array([[4.0, 5.0, 6.0]], dtype=np.float32), + gravity=np.array([[0.1, 0.2, 0.9]], dtype=np.float32), + dof_pos=np.zeros((1, 29), dtype=np.float32), + dof_vel=np.array([np.arange(7.0, 36.0, dtype=np.float32)], dtype=np.float32), + ), + ) + + +@pytest.mark.parametrize("case", _G1_OWNER_CASES, ids=[case["id"] for case in _G1_OWNER_CASES]) +def test_g1_owner_yaml_regression_contract(case: dict[str, Any]): + from unilab.envs.locomotion.g1.joystick import G1WalkEnv + + cfg = _compose_cfg(case["config_group"], case["overrides"]) + full_env_cfg = _materialize_env_cfg(cfg) + env = _build_probe_env(cfg) + env_cfg_override = BackendAdapter( + cfg, root_dir=ROOT_DIR, algo_name=cfg.algo.algo if "algo" in cfg.algo else None + ).build_task_env_cfg_override() + + assert cfg.training.task_name == case["task_name"] + assert cfg.training.sim_backend == case["backend"] + assert full_env_cfg.control_config.action_scale == pytest.approx(case["action_scale"]) + assert full_env_cfg.curriculum.enabled is case["curriculum_enabled"] + assert env._uses_walk_observation_profile() is (case["profile"] == "walk") + assert ( + registry._envs[cfg.training.task_name].env_cls_dict[cfg.training.sim_backend] is G1WalkEnv + ) + + if "model_suffix" in case: + assert str(full_env_cfg.model_file).endswith(case["model_suffix"]) + + reward_config = OmegaConf.to_container(cfg.reward, resolve=True) + assert env_cfg_override["reward_config"] == reward_config + env_override = cast(dict[str, Any], OmegaConf.to_container(cfg.env, resolve=True)) + for key, value in env_override.items(): + assert env_cfg_override[key] == value + + env._reward_fns = {} + env._init_reward_functions() + for reward_name in cfg.reward.scales.keys(): + assert reward_name in env._reward_fns + + +@pytest.mark.parametrize("case", _G1_OWNER_CASES, ids=[case["id"] for case in _G1_OWNER_CASES]) +def test_g1_owner_yaml_observation_profiles_match_expected_family(case: dict[str, Any]): + cfg = _compose_cfg(case["config_group"], case["overrides"]) + obs = _compute_probe_obs(cfg) + + if case["profile"] == "legacy": + np.testing.assert_allclose(obs["obs"][:, :3], [[104.0, 105.0, 106.0]]) + np.testing.assert_allclose(obs["obs"][:, 35:37], [[107.0, 108.0]]) + np.testing.assert_allclose(obs["critic"][:, :3], [[4.0, 5.0, 6.0]]) + np.testing.assert_allclose(obs["critic"][:, 35:37], [[7.0, 8.0]]) + np.testing.assert_allclose(obs["critic"][:, 98:101], [[1.0, 2.0, 3.0]]) + else: + np.testing.assert_allclose(obs["obs"][:, :3], [[26.0, 26.25, 26.5]]) + np.testing.assert_allclose(obs["obs"][:, 35:37], [[5.35, 5.4]]) + np.testing.assert_allclose(obs["critic"][:, :3], [[1.0, 1.25, 1.5]]) + np.testing.assert_allclose(obs["critic"][:, 35:37], [[0.35, 0.4]]) + np.testing.assert_allclose(obs["critic"][:, 98:101], [[2.0, 4.0, 6.0]]) + + +def test_g1_observation_profile_selection_prefers_reward_family_over_curriculum_flag(): + from unilab.envs.locomotion.g1.joystick import G1WalkEnv + + env = cast(Any, object.__new__(G1WalkEnv)) + + env._cfg = cast( + Any, + type( + "Cfg", + (), + {"curriculum": type("Curriculum", (), {"enabled": True})(), "reward_config": None}, + )(), + ) + env._reward_cfg = cast( + Any, + type("RewardCfg", (), {"scales": {"orientation": -2.5, "ang_vel_xy": -0.2}})(), + ) + assert env._uses_walk_observation_profile() is False + + env._cfg = cast( + Any, + type( + "Cfg", + (), + {"curriculum": type("Curriculum", (), {"enabled": False})(), "reward_config": None}, + )(), + ) + env._reward_cfg = cast( + Any, + type( + "RewardCfg", + (), + {"scales": {"penalty_orientation": -10.0, "penalty_ang_vel_xy": -1.0, "alive": 10.0}}, + )(), + ) + assert env._uses_walk_observation_profile() is True + + +def test_g1_walk_tasks_are_registered(): + ensure_registries() + + assert registry.contains("G1WalkFlat") + assert registry.contains("G1WalkRough") diff --git a/tests/envs/locomotion/g1/test_symmetry_contract.py b/tests/envs/locomotion/g1/test_symmetry_contract.py index c40578b19..857a95d29 100644 --- a/tests/envs/locomotion/g1/test_symmetry_contract.py +++ b/tests/envs/locomotion/g1/test_symmetry_contract.py @@ -6,14 +6,14 @@ import torch from unilab.base import registry -from unilab.envs.locomotion.g1.joystick_sac import RewardConfigSAC +from unilab.envs.locomotion.g1.joystick import G1WalkRewardConfig from unilab.utils.algo_utils import ensure_registries pytest.importorskip("mujoco", reason="mujoco is required for G1 symmetry contract tests") -def _reward_config() -> RewardConfigSAC: - return RewardConfigSAC( +def _reward_config() -> G1WalkRewardConfig: + return G1WalkRewardConfig( scales={"tracking_lin_vel": 2.0, "alive": 10.0}, tracking_sigma=0.25, base_height_target=0.754, diff --git a/tests/envs/test_env_configs.py b/tests/envs/test_env_configs.py index 2c2d02be8..9c8e4ccf4 100644 --- a/tests/envs/test_env_configs.py +++ b/tests/envs/test_env_configs.py @@ -74,14 +74,14 @@ def blocked_import(name, globals=None, locals=None, fromlist=(), level=0): assert result.returncode == 0, result.stderr or result.stdout -def test_g1_joystick_flat_ppo_cfg_obs_groups_spec(): - """G1JoystickPPO must declare obs_groups_spec with actor and critic groups.""" - from unilab.envs.locomotion.g1.joystick import G1JoystickPPOCfg, RewardConfigPPO +def test_g1_walk_env_cfg_obs_groups_spec(): + """G1WalkEnv must declare obs_groups_spec with actor and critic groups.""" + from unilab.envs.locomotion.g1.joystick import G1WalkEnvCfg, G1WalkLegacyRewardConfig - cfg = G1JoystickPPOCfg() + cfg = G1WalkEnvCfg() assert not hasattr(cfg, "obs_config"), "obs_config should have been removed" - reward_cfg = RewardConfigPPO( + reward_cfg = G1WalkLegacyRewardConfig( scales={"feet_phase": 1.0}, tracking_sigma=0.25, gait_frequency=1.5, @@ -97,7 +97,7 @@ def test_g1_joystick_flat_ppo_cfg_obs_groups_spec(): def test_g1_walk_flat_cfg_no_obs_config(): """G1WalkFlatCfg should no longer have obs_config after dict obs refactor.""" - from unilab.envs.locomotion.g1.joystick_sac import G1WalkFlatCfg + from unilab.envs.locomotion.g1.joystick import G1WalkFlatCfg cfg = G1WalkFlatCfg() assert not hasattr(cfg, "obs_config"), ( @@ -106,7 +106,7 @@ def test_g1_walk_flat_cfg_no_obs_config(): def test_g1_walk_flat_cfg_has_domain_rand_for_motrix(): - from unilab.envs.locomotion.g1.joystick_sac import G1WalkFlatCfg + from unilab.envs.locomotion.g1.joystick import G1WalkFlatCfg cfg = G1WalkFlatCfg() assert hasattr(cfg, "domain_rand") @@ -117,27 +117,123 @@ def test_g1_walk_flat_cfg_has_domain_rand_for_motrix(): assert cfg.domain_rand.push_robots is False -def test_g1_joystick_flat_ppo_obs_groups_spec_dims(): +def test_g1_walk_flat_cfg_defaults_match_walk_profile(): + from unilab.envs.locomotion.g1.joystick import G1WalkFlatCfg + + cfg = G1WalkFlatCfg() + assert not hasattr(cfg, "obs_profile") + assert cfg.curriculum.enabled is True + + +def test_g1_walk_tasks_register_to_algorithm_agnostic_env_base(): + from unilab.base import registry + from unilab.envs.locomotion.g1.joystick import G1WalkEnv, G1WalkRewardConfig + + env = cast( + Any, + registry.make( + "G1WalkFlat", + num_envs=1, + sim_backend="mujoco", + env_cfg_override={ + "reward_config": G1WalkRewardConfig( + scales={"tracking_lin_vel": 2.0, "alive": 10.0}, + tracking_sigma=0.25, + base_height_target=0.754, + min_base_height=0.3, + max_tilt_deg=65.0, + gait_frequency=1.5, + feet_phase_swing_height=0.09, + feet_phase_tracking_sigma=0.04, + close_feet_threshold=0.15, + pose_weights=[0.01] * 29, + ) + }, + ), + ) + try: + assert env.__class__ is G1WalkEnv + finally: + env.close() + + +def test_g1_walk_flat_observation_construction_is_hardcoded_for_legacy_and_walk_modes(): + from unilab.envs.locomotion.g1.joystick import G1WalkEnv + + class NoiseCfg: + level = 0.0 + scale_gyro = 0.0 + scale_gravity = 0.0 + scale_joint_angle = 0.0 + scale_joint_vel = 0.0 + scale_linvel = 0.0 + + def compute_obs(curriculum_enabled: bool) -> dict[str, np.ndarray]: + env = cast(Any, object.__new__(G1WalkEnv)) + env._num_envs = 1 + env.default_angles = np.array([[0.5, -0.5]], dtype=np.float32) + env._cfg = type( + "Cfg", + (), + { + "noise_config": NoiseCfg(), + "curriculum": type("Curriculum", (), {"enabled": curriculum_enabled})(), + }, + )() + env._obs_noise = lambda data, scale: data + 100.0 + info = { + "commands": np.array([[0.7, 0.0, 0.2]], dtype=np.float32), + "current_actions": np.array([[0.1, -0.2]], dtype=np.float32), + "gait_phase": np.array([[0.3, 3.4]], dtype=np.float32), + } + return cast( + dict[str, np.ndarray], + env._compute_obs( + info, + linvel=np.array([[1.0, 2.0, 3.0]], dtype=np.float32), + gyro=np.array([[4.0, 5.0, 6.0]], dtype=np.float32), + gravity=np.array([[0.1, 0.2, 0.9]], dtype=np.float32), + dof_pos=np.array([[0.6, -0.3]], dtype=np.float32), + dof_vel=np.array([[7.0, 8.0]], dtype=np.float32), + ), + ) + + legacy = compute_obs(curriculum_enabled=False) + walk = compute_obs(curriculum_enabled=True) + + np.testing.assert_allclose(legacy["obs"][:, :3], [[104.0, 105.0, 106.0]]) + np.testing.assert_allclose(legacy["obs"][:, 8:10], [[107.0, 108.0]]) + np.testing.assert_allclose(legacy["critic"][:, :3], [[4.0, 5.0, 6.0]]) + np.testing.assert_allclose(legacy["critic"][:, 17:20], [[1.0, 2.0, 3.0]]) + + np.testing.assert_allclose(walk["obs"][:, :3], [[26.0, 26.25, 26.5]]) + np.testing.assert_allclose(walk["obs"][:, 8:10], [[5.35, 5.4]]) + np.testing.assert_allclose(walk["critic"][:, :3], [[1.0, 1.25, 1.5]]) + np.testing.assert_allclose(walk["critic"][:, 8:10], [[0.35, 0.4]]) + np.testing.assert_allclose(walk["critic"][:, 17:20], [[2.0, 4.0, 6.0]]) + + +def test_g1_walk_env_obs_groups_spec_dims(): """obs_groups_spec total dim must match what _compute_obs actually produces. - G1JoystickPPO._compute_obs outputs (G1 has 29 DoF): + G1WalkEnv._compute_obs outputs (G1 has 29 DoF): actor: gyro(3) + gravity(3) + diff(29) + dof_vel(29) + last_actions(29) + command(3) + gait_phase(2) = 98 critic: actor(98) + linvel(3) = 101 """ - from unilab.envs.locomotion.g1.joystick import G1JoystickPPO + from unilab.envs.locomotion.g1.joystick import G1WalkEnv # obs_groups_spec is a @property; access via descriptor protocol - spec = G1JoystickPPO.obs_groups_spec.fget(None) # type: ignore[union-attr] + spec = G1WalkEnv.obs_groups_spec.fget(None) # type: ignore[union-attr] assert spec is not None assert spec["obs"] == 98 assert spec["critic"] == 101 -def test_g1_joystick_flat_ppo_reward_dispatch_restores_motrix_terms(): - from unilab.envs.locomotion.g1.joystick import G1JoystickPPO +def test_g1_walk_env_reward_dispatch_restores_motrix_terms(): + from unilab.envs.locomotion.g1.joystick import G1WalkEnv - env = cast(Any, object.__new__(G1JoystickPPO)) + env = cast(Any, object.__new__(G1WalkEnv)) env._reward_fns = {} env._init_reward_functions() @@ -148,9 +244,9 @@ def test_g1_joystick_flat_ppo_reward_dispatch_restores_motrix_terms(): assert "feet_double_stance" in env._reward_fns -def test_g1_joystick_flat_ppo_feet_phase_reward_is_gated_by_forward_speed(): +def test_g1_walk_env_feet_phase_reward_is_gated_by_forward_speed(): from unilab.envs.locomotion.common.rewards import RewardContext - from unilab.envs.locomotion.g1.joystick import G1JoystickPPO + from unilab.envs.locomotion.g1.joystick import G1WalkEnv class FakeBackend: def get_sensor_data(self, name: str) -> np.ndarray: @@ -160,7 +256,7 @@ def get_sensor_data(self, name: str) -> np.ndarray: return np.array([[0.0, 0.0, 0.0], [0.0, 0.0, 0.0]], dtype=np.float32) raise KeyError(name) - env = cast(Any, object.__new__(G1JoystickPPO)) + env = cast(Any, object.__new__(G1WalkEnv)) env._backend = FakeBackend() env._num_envs = 2 env._reward_cfg = type( @@ -187,7 +283,7 @@ def get_sensor_data(self, name: str) -> np.ndarray: assert reward[1] > 0.0 -def test_g1_joystick_flat_assets_define_contact_sensors_for_gait_rewards(): +def test_g1_walk_flat_assets_define_contact_sensors_for_gait_rewards(): repo_root = Path(__file__).parents[2] scene_text = ( repo_root / "src" / "unilab" / "assets" / "robots" / "g1" / "scene_flat.xml" @@ -510,7 +606,6 @@ def test_g1_motion_tracking_clip_end_does_not_override_true_termination(): _STANDARD_ENVS = [ "Go1JoystickFlat", "Go2JoystickFlat", - "G1JoystickFlat", "G1WalkFlat", "G1WalkRough", "AllegroInhandRotation", diff --git a/tests/scripts/test_train_script_configs.py b/tests/scripts/test_train_script_configs.py index b4bae5d1c..f114b931e 100644 --- a/tests/scripts/test_train_script_configs.py +++ b/tests/scripts/test_train_script_configs.py @@ -32,7 +32,7 @@ def _mlx_runtime_usable() -> bool: [ "go1_joystick_flat/mujoco", "go2_joystick_flat/mujoco", - "g1_joystick_flat/mujoco", + "g1_walk_flat/mujoco", "g1_motion_tracking/mujoco", "g1_flip_tracking/mujoco", ], diff --git a/tests/scripts/test_train_scripts.py b/tests/scripts/test_train_scripts.py index 6698957fd..a34d49cb7 100644 --- a/tests/scripts/test_train_scripts.py +++ b/tests/scripts/test_train_scripts.py @@ -222,7 +222,7 @@ def test_ppo_g1_resolved_algo_matches_old_motrix_behavior(): For this migration we align with the final UniLab1 Motrix runtime. """ - cfg = _ppo_cfg(["task=g1_joystick_flat/motrix"]) + cfg = _ppo_cfg(["task=g1_walk_flat/motrix"]) assert cfg.algo.max_iterations == 220 assert cfg.algo.empirical_normalization is True @@ -233,7 +233,7 @@ def test_ppo_g1_resolved_algo_matches_old_motrix_behavior(): def test_ppo_g1_mujoco_base_hyperparams_remain_separate(): - cfg = _ppo_cfg(["task=g1_joystick_flat/mujoco"]) + cfg = _ppo_cfg(["task=g1_walk_flat/mujoco"]) assert cfg.algo.max_iterations == 220 assert cfg.algo.empirical_normalization is False @@ -241,7 +241,7 @@ def test_ppo_g1_mujoco_base_hyperparams_remain_separate(): def test_ppo_g1_env_preset_has_env_overrides(): - cfg = _ppo_cfg(["task=g1_joystick_flat/motrix"]) + cfg = _ppo_cfg(["task=g1_walk_flat/motrix"]) assert cfg.env.iterations == 3 assert cfg.env.control_config.action_scale == pytest.approx(0.5) @@ -285,7 +285,7 @@ def test_build_ppo_env_cfg_override_g1_motrix( monkeypatch: pytest.MonkeyPatch, ): mod = _train_rsl_rl(monkeypatch) - cfg = _ppo_cfg(["task=g1_joystick_flat/motrix"]) + cfg = _ppo_cfg(["task=g1_walk_flat/motrix"]) env_cfg_override = mod.build_ppo_env_cfg_override(cfg) @@ -410,7 +410,7 @@ def test_ppo_cli_algo_override_wins_over_base( monkeypatch: pytest.MonkeyPatch, ): """CLI override takes precedence over base task algo values via Hydra compose.""" - cfg = _ppo_cfg(["task=g1_joystick_flat/motrix", "algo.max_iterations=1"]) + cfg = _ppo_cfg(["task=g1_walk_flat/motrix", "algo.max_iterations=1"]) assert cfg.algo.max_iterations == 1 # Other base values remain intact @@ -1260,14 +1260,14 @@ def test_offpolicy_td3_hydra_default_algo_log_name(): def test_offpolicy_flashsac_hydra_algo_log_name(): - cfg = _offpolicy_cfg(["algo=flashsac", "task=flashsac/g1_joystick_flat/mujoco"]) + cfg = _offpolicy_cfg(["algo=flashsac", "task=flashsac/g1_walk_flat_amp/mujoco"]) assert cfg.algo.algo_log_name == "flash_sac" assert cfg.algo.load_run == "-1" -def test_offpolicy_flashsac_g1_joystick_flat_task_composes() -> None: - cfg = _offpolicy_cfg(["algo=flashsac", "task=flashsac/g1_joystick_flat/mujoco"]) - assert cfg.training.task_name == "G1JoystickFlat" +def test_offpolicy_flashsac_g1_walk_flat_amp_task_composes() -> None: + cfg = _offpolicy_cfg(["algo=flashsac", "task=flashsac/g1_walk_flat_amp/mujoco"]) + assert cfg.training.task_name == "G1WalkFlat" assert cfg.training.sim_backend == "mujoco" @@ -1282,7 +1282,7 @@ def test_offpolicy_flashsac_rejects_multi_gpu(): cfg = _offpolicy_cfg( [ "algo=flashsac", - "task=flashsac/g1_joystick_flat/mujoco", + "task=flashsac/g1_walk_flat_amp/mujoco", "training.num_gpus=2", ] ) From 41ecaa4f7609a75ebc0b89320c9e739d14057774 Mon Sep 17 00:00:00 2001 From: tatp-yf Date: Wed, 22 Apr 2026 15:04:48 +0800 Subject: [PATCH 2/4] fix: unify critic_obs_dim naming and remove fallback compat behavior MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - split_obs_dict / get_obs_dims: when no critic group, critic = obs (never return None or 0) - All learners (TD3/SAC/FlashSAC): critic_obs_dim is now a required parameter, no default-0-then-fallback pattern - update_critic / update_actor: use data["critic"] directly instead of data.get("critic") with fallback to actor obs - OffPolicyRunner: rename critic_dim → critic_obs_dim for consistency - Collectors: critic_np is always a concrete array, no None branches - FastTD3Runner: pass env_cfg_override to get_env_dims and critic_obs_dim to learner (fixes td3/g1_walk_flat config load) --- src/unilab/algos/torch/appo/worker.py | 10 ++-- src/unilab/algos/torch/fast_sac/learner.py | 60 +++++++------------ src/unilab/algos/torch/fast_td3/learner.py | 14 ++--- src/unilab/algos/torch/fast_td3/runner.py | 3 +- src/unilab/algos/torch/flash_sac/learner.py | 16 ++--- src/unilab/algos/torch/offpolicy/runner.py | 4 +- src/unilab/algos/torch/offpolicy/worker.py | 10 ++-- src/unilab/utils/final_observation.py | 6 +- src/unilab/utils/obs_utils.py | 18 ++++-- tests/algos/test_appo_worker.py | 1 + .../algos/test_fast_sac_symmetry_contract.py | 1 + tests/algos/test_offpolicy_runner.py | 2 + tests/utils/test_obs_utils.py | 5 +- 13 files changed, 65 insertions(+), 85 deletions(-) diff --git a/src/unilab/algos/torch/appo/worker.py b/src/unilab/algos/torch/appo/worker.py index 17d41ba3a..a9a645d05 100644 --- a/src/unilab/algos/torch/appo/worker.py +++ b/src/unilab/algos/torch/appo/worker.py @@ -26,7 +26,7 @@ def compute_timeout_bootstrap_correction( gamma: float, timeout_mask: np.ndarray, final_obs: np.ndarray, - final_critic: np.ndarray | None = None, + final_critic: np.ndarray, ) -> np.ndarray: """Compute gamma * V(final_observation) for current timeout envs.""" corrections = np.zeros(timeout_mask.shape, dtype=np.float32) @@ -35,7 +35,7 @@ def compute_timeout_bootstrap_correction( from tensordict import TensorDict - critic_input_np = final_critic if final_critic is not None else final_obs + critic_input_np = final_critic critic_input = torch.from_numpy(critic_input_np[timeout_mask]).to(collector_device) critic_td = TensorDict( {"policy": critic_input}, @@ -172,8 +172,7 @@ def to_float32_np(x): obs_np, critic_np = split_obs_dict(obs_out) obs_np = to_float32_np(obs_np) - if critic_np is not None: - critic_np = to_float32_np(critic_np) + critic_np = to_float32_np(critic_np) # Pre-allocate obs TensorDict once; update in-place each step to avoid # repeated TensorDict construction overhead in the hot loop. @@ -243,8 +242,7 @@ def to_float32_np(x): next_actor_obs_np, next_critic_np = split_obs_dict(next_obs_raw) 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) + next_critic_np = to_float32_np(next_critic_np) terminal_contract = resolve_terminal_observation_contract( next_obs_batch_size=next_actor_obs_np.shape[0], final_observation=getattr(state, "final_observation", None), diff --git a/src/unilab/algos/torch/fast_sac/learner.py b/src/unilab/algos/torch/fast_sac/learner.py index 0c98abb35..130483c02 100644 --- a/src/unilab/algos/torch/fast_sac/learner.py +++ b/src/unilab/algos/torch/fast_sac/learner.py @@ -353,6 +353,7 @@ def __init__( self, obs_dim: int, action_dim: int, + critic_obs_dim: int, device: str = "cpu", # Hyperparameters aligned with holosoma gamma: float = 0.97, @@ -379,7 +380,6 @@ def __init__( use_amp: bool = False, symmetry_augmentation: SymmetryAugmentation | None = None, world_size: int = 1, - critic_obs_dim: int = 0, ): self.device = device self.gamma = gamma @@ -402,9 +402,8 @@ def __init__( device=device, ) - critic_net_obs_dim = critic_obs_dim if critic_obs_dim > 0 else obs_dim self.qnet = SACCritic( - obs_dim=critic_net_obs_dim, + obs_dim=critic_obs_dim, action_dim=action_dim, num_atoms=num_atoms, v_min=v_min, @@ -417,7 +416,7 @@ def __init__( # Target critic self.qnet_target = SACCritic( - obs_dim=critic_net_obs_dim, + obs_dim=critic_obs_dim, action_dim=action_dim, num_atoms=num_atoms, v_min=v_min, @@ -496,17 +495,14 @@ def _reduce_gradients(self, model: nn.Module) -> None: def update_critic(self, batch: Dict[str, torch.Tensor]) -> Dict[str, float]: """One critic update step.""" obs = batch["obs"] - critic = batch.get("critic", None) + critic_obs = batch["critic"] actions = batch["actions"] rewards = batch["rewards"] next_obs = batch["next_obs"] - next_critic = batch.get("next_critic", None) + critic_next_obs = batch["next_critic"] dones = batch["dones"] truncated = batch.get("truncated") - critic_obs = critic if critic is not None else obs - critic_next_obs = next_critic if next_critic is not None else next_obs - # Apply symmetry augmentation if self.use_symmetry: orig_actions = actions @@ -517,24 +513,16 @@ def update_critic(self, batch: Dict[str, torch.Tensor]) -> Dict[str, float]: next_obs, orig_actions, obs_group="obs" ) - if critic is not None: - assert next_critic is not None - critic_base_aug, _ = self.symmetry.augment_obs_and_actions( - critic, - orig_actions, - obs_group="critic", - ) - critic_next_base_aug, _ = self.symmetry.augment_obs_and_actions( - next_critic, - orig_actions, - obs_group="critic", - ) - else: - critic_base_aug = obs - critic_next_base_aug = next_obs - - critic_obs = critic_base_aug - critic_next_obs = critic_next_base_aug + critic_obs, _ = self.symmetry.augment_obs_and_actions( + critic_obs, + orig_actions, + obs_group="critic", + ) + critic_next_obs, _ = self.symmetry.augment_obs_and_actions( + critic_next_obs, + orig_actions, + obs_group="critic", + ) # Double the batch size for other tensors rewards = rewards.repeat(2) @@ -624,24 +612,16 @@ def update_critic(self, batch: Dict[str, torch.Tensor]) -> Dict[str, float]: def update_actor(self, batch: Dict[str, torch.Tensor]) -> Dict[str, float]: """One actor update step.""" obs = batch["obs"] - critic = batch.get("critic", None) - - critic_obs = critic if critic is not None else obs + critic_obs = batch["critic"] # Apply symmetry augmentation if self.use_symmetry: assert self.symmetry is not None obs = torch.cat([obs, self.symmetry.mirror_obs(obs, obs_group="obs")], dim=0) - - if critic is not None: - critic_base_aug = torch.cat( - [critic, self.symmetry.mirror_obs(critic, obs_group="critic")], - dim=0, - ) - else: - critic_base_aug = obs - - critic_obs = critic_base_aug + critic_obs = torch.cat( + [critic_obs, self.symmetry.mirror_obs(critic_obs, obs_group="critic")], + dim=0, + ) with torch.amp.autocast("cuda", enabled=self.use_amp): # pyright: ignore[reportPrivateImportUsage] actions, log_probs, log_std = self.actor.get_actions_and_log_probs(obs) diff --git a/src/unilab/algos/torch/fast_td3/learner.py b/src/unilab/algos/torch/fast_td3/learner.py index db8e67447..3c4a7b268 100644 --- a/src/unilab/algos/torch/fast_td3/learner.py +++ b/src/unilab/algos/torch/fast_td3/learner.py @@ -138,7 +138,7 @@ def __init__( self, obs_dim: int, action_dim: int, - critic_obs_dim: int = 0, + critic_obs_dim: int, num_envs: int = 1024, device: str = "cpu", # Hyperparameters from reference @@ -171,7 +171,7 @@ def __init__( self.noise_clip = noise_clip self.policy_frequency = policy_frequency self.use_cdq = use_cdq - self.critic_obs_dim = critic_obs_dim if critic_obs_dim > 0 else obs_dim + self.critic_obs_dim = critic_obs_dim torch_device = torch.device(device) @@ -250,17 +250,14 @@ def normalize_obs(self, obs: torch.Tensor, update: bool = False) -> torch.Tensor def update_critic(self, data: Dict[str, torch.Tensor]) -> Dict[str, float]: """One critic update step.""" observations = self.normalize_obs(data["obs"], update=True) - critic_obs = data.get("critic", None) + critic_obs = data["critic"] actions = data["actions"] rewards = data["rewards"] next_observations = self.normalize_obs(data["next_obs"], update=False) - next_critic_obs = data.get("next_critic", None) + critic_next_obs = data["next_critic"] dones = data["dones"].bool() truncations = data["truncated"].bool() - critic_obs = critic_obs if critic_obs is not None else observations - critic_next_obs = next_critic_obs if next_critic_obs is not None else next_observations - bootstrap = (truncations | ~dones).float() discount = torch.full_like(rewards, self.gamma) @@ -325,8 +322,9 @@ def update_critic(self, data: Dict[str, torch.Tensor]) -> Dict[str, float]: def update_actor(self, data: Dict[str, torch.Tensor]) -> Dict[str, float]: """One actor update step.""" observations = self.normalize_obs(data["obs"], update=False) + critic_obs = data["critic"] - qf1, qf2 = self.qnet(observations, self.actor(observations)) + qf1, qf2 = self.qnet(critic_obs, self.actor(observations)) qf1_value = self.qnet.get_value(F.softmax(qf1, dim=1)) qf2_value = self.qnet.get_value(F.softmax(qf2, dim=1)) diff --git a/src/unilab/algos/torch/fast_td3/runner.py b/src/unilab/algos/torch/fast_td3/runner.py index 39d6797e3..77106d74e 100644 --- a/src/unilab/algos/torch/fast_td3/runner.py +++ b/src/unilab/algos/torch/fast_td3/runner.py @@ -42,10 +42,11 @@ def __init__( obs_normalization: bool = True, sim_backend: str = "mujoco", ): - obs_dim, action_dim, _ = get_env_dims(env_name, sim_backend) + obs_dim, action_dim, critic_obs_dim = get_env_dims(env_name, sim_backend, env_cfg_override=env_cfg_override) learner = FastTD3Learner( obs_dim=obs_dim, action_dim=action_dim, + critic_obs_dim=critic_obs_dim, num_envs=num_envs, device=device or get_default_device(), gamma=gamma, diff --git a/src/unilab/algos/torch/flash_sac/learner.py b/src/unilab/algos/torch/flash_sac/learner.py index 025f15197..aee22dc49 100644 --- a/src/unilab/algos/torch/flash_sac/learner.py +++ b/src/unilab/algos/torch/flash_sac/learner.py @@ -141,7 +141,7 @@ def __init__( self, obs_dim: int, action_dim: int, - critic_obs_dim: int = 0, + critic_obs_dim: int, device: str = "cpu", gamma: float = 0.99, tau: float = 0.01, @@ -178,7 +178,7 @@ def __init__( self.n_step = n_step self.actor_bc_alpha = actor_bc_alpha self.obs_dim = obs_dim - self.critic_obs_dim = critic_obs_dim if critic_obs_dim > 0 else obs_dim + self.critic_obs_dim = critic_obs_dim self.action_dim = action_dim self.update_count = 0 self.use_amp = bool(use_amp and self.device.type == "cuda") @@ -283,15 +283,11 @@ def update_critic(self, batch: dict[str, torch.Tensor]) -> dict[str, float]: next_obs = batch["next_obs"].to(self.device) terminated = batch["dones"].to(self.device) truncated = batch.get("truncated", torch.zeros_like(terminated)).to(self.device) - critic_obs = batch.get("critic") - next_critic_obs = batch.get("next_critic") - critic_obs = critic_obs.to(self.device) if critic_obs is not None else None - next_critic_obs = next_critic_obs.to(self.device) if next_critic_obs is not None else None + critic_obs = batch["critic"].to(self.device) + critic_next_obs = batch["next_critic"].to(self.device) obs = self._maybe_normalize_obs(obs, update=True) next_obs = self._maybe_normalize_obs(next_obs, update=False) - critic_obs = critic_obs if critic_obs is not None else obs - critic_next_obs = next_critic_obs if next_critic_obs is not None else next_obs if self.reward_normalizer is not None: rewards = self.reward_normalizer.normalize(rewards) @@ -349,12 +345,10 @@ def update_actor(self, batch: dict[str, torch.Tensor]) -> dict[str, float]: obs = batch["obs"].to(self.device) next_obs = batch["next_obs"].to(self.device) expert_actions = batch["actions"].to(self.device) - critic_obs = batch.get("critic") - critic_obs = critic_obs.to(self.device) if critic_obs is not None else None + critic_obs = batch["critic"].to(self.device) obs = self._maybe_normalize_obs(obs, update=False) next_obs = self._maybe_normalize_obs(next_obs, update=False) - critic_obs = critic_obs if critic_obs is not None else obs obs_all = torch.cat([obs, next_obs], dim=0) diff --git a/src/unilab/algos/torch/offpolicy/runner.py b/src/unilab/algos/torch/offpolicy/runner.py index 51c579105..2720c4f51 100644 --- a/src/unilab/algos/torch/offpolicy/runner.py +++ b/src/unilab/algos/torch/offpolicy/runner.py @@ -91,7 +91,7 @@ def __init__( self.obs_normalization = obs_normalization self.actor_kwargs = actor_kwargs or {} - self.obs_dim, self.action_dim, self.critic_dim = get_env_dims( + self.obs_dim, self.action_dim, self.critic_obs_dim = get_env_dims( self.env_name, sim_backend, env_cfg_override ) @@ -177,7 +177,7 @@ def learn( obs_dim=self.obs_dim, action_dim=self.action_dim, device=self.device, - critic_dim=self.critic_dim, + critic_dim=self.critic_obs_dim, ) self._shared_resources.append(replay_buffer) diff --git a/src/unilab/algos/torch/offpolicy/worker.py b/src/unilab/algos/torch/offpolicy/worker.py index 12150e394..818ff65c1 100644 --- a/src/unilab/algos/torch/offpolicy/worker.py +++ b/src/unilab/algos/torch/offpolicy/worker.py @@ -206,8 +206,7 @@ def _run_collector( state = env.step(actions_np) obs_np, critic_np = split_obs_dict(state.obs) obs_np = np.asarray(obs_np, dtype=np.float32) - if critic_np is not None: - critic_np = np.asarray(critic_np, dtype=np.float32) + critic_np = np.asarray(critic_np, dtype=np.float32) prev_dones_np = np.zeros(num_envs, dtype=np.float32) import time as _time @@ -266,8 +265,7 @@ def _run_collector( # Extract data as numpy next_obs_np, next_critic_np = split_obs_dict(state.obs) next_obs_np = np.asarray(next_obs_np, dtype=np.float32) - if next_critic_np is not None: - next_critic_np = np.asarray(next_critic_np, dtype=np.float32) + next_critic_np = np.asarray(next_critic_np, dtype=np.float32) rewards_np = np.asarray(state.reward, dtype=np.float32).ravel() terminated_np = ( @@ -312,8 +310,8 @@ def _run_collector( if terminal_contract.terminal_obs is not None else None ), - critic=torch.from_numpy(critic_np) if critic_np is not None else None, - next_critic=torch.from_numpy(next_critic_np) if next_critic_np is not None else None, + critic=torch.from_numpy(critic_np), + next_critic=torch.from_numpy(next_critic_np), terminal_next_critic=( torch.from_numpy(terminal_contract.terminal_critic) if terminal_contract.terminal_critic is not None diff --git a/src/unilab/utils/final_observation.py b/src/unilab/utils/final_observation.py index 69e1f2291..08937f593 100644 --- a/src/unilab/utils/final_observation.py +++ b/src/unilab/utils/final_observation.py @@ -5,9 +5,6 @@ import numpy as np -from unilab.utils.obs_utils import split_obs_dict - - @dataclass(frozen=True) class TransitionBootstrapContract: actor_next_obs: np.ndarray @@ -115,7 +112,8 @@ def resolve_terminal_observation_contract( terminal_obs: np.ndarray | None = None terminal_critic: np.ndarray | None = None if np.any(terminal_mask) and isinstance(resolved_final_observation, dict): - terminal_obs, terminal_critic = split_obs_dict(resolved_final_observation) + terminal_obs = resolved_final_observation.get("obs") + terminal_critic = resolved_final_observation.get("critic") timeout_terminal_mask = terminal_mask if truncated is not None: diff --git a/src/unilab/utils/obs_utils.py b/src/unilab/utils/obs_utils.py index aab7a43d8..3b17f9cfa 100644 --- a/src/unilab/utils/obs_utils.py +++ b/src/unilab/utils/obs_utils.py @@ -13,14 +13,22 @@ def flatten_policy_obs_dict(obs: dict[str, np.ndarray]) -> np.ndarray: return obs["obs"] -def split_obs_dict(obs: dict[str, np.ndarray]) -> tuple[np.ndarray, np.ndarray | None]: - """Split observation dict into (actor_obs, critic_obs).""" - return obs["obs"], obs.get("critic") +def split_obs_dict(obs: dict[str, np.ndarray]) -> tuple[np.ndarray, np.ndarray]: + """Split observation dict into (actor_obs, critic_obs). + + When no separate critic group exists, critic_obs == actor_obs. + """ + actor = obs["obs"] + return actor, obs.get("critic", actor) def get_obs_dims(obs_groups_spec: dict[str, int]) -> tuple[int, int]: - """Extract (actor_obs_dim, critic_obs_dim) from obs_groups_spec.""" - return obs_groups_spec.get("obs", 0), obs_groups_spec.get("critic", 0) + """Extract (actor_obs_dim, critic_obs_dim) from obs_groups_spec. + + When no separate critic group exists, critic_obs_dim == actor_obs_dim. + """ + obs_dim = obs_groups_spec.get("obs", 0) + return obs_dim, obs_groups_spec.get("critic", obs_dim) def get_critic_base_dim(obs_groups_spec: dict[str, int]) -> int: diff --git a/tests/algos/test_appo_worker.py b/tests/algos/test_appo_worker.py index d892cbadd..962e42ef9 100644 --- a/tests/algos/test_appo_worker.py +++ b/tests/algos/test_appo_worker.py @@ -19,6 +19,7 @@ def test_compute_timeout_bootstrap_correction_uses_final_observation_value(): gamma=0.5, timeout_mask=np.array([True, False]), final_obs=np.array([[2.0, 3.0], [9.0, 9.0]], dtype=np.float32), + final_critic=np.array([[2.0, 3.0], [9.0, 9.0]], dtype=np.float32), ) np.testing.assert_allclose(correction, np.array([2.5, 0.0], dtype=np.float32)) diff --git a/tests/algos/test_fast_sac_symmetry_contract.py b/tests/algos/test_fast_sac_symmetry_contract.py index ad7498fae..0dfb1ad76 100644 --- a/tests/algos/test_fast_sac_symmetry_contract.py +++ b/tests/algos/test_fast_sac_symmetry_contract.py @@ -115,6 +115,7 @@ def test_fast_sac_learner_rejects_symmetry_without_augmentation(): FastSACLearner( obs_dim=4, action_dim=2, + critic_obs_dim=4, device="cpu", use_symmetry=True, ) diff --git a/tests/algos/test_offpolicy_runner.py b/tests/algos/test_offpolicy_runner.py index a47ac1798..81dcc6bef 100644 --- a/tests/algos/test_offpolicy_runner.py +++ b/tests/algos/test_offpolicy_runner.py @@ -27,6 +27,7 @@ def _make_sac_runner(env_name: str) -> OffPolicyRunner: learner = FastSACLearner( obs_dim=obs_dim, action_dim=action_dim, + critic_obs_dim=obs_dim, device="cpu", actor_hidden_dim=cfg.get("actor_hidden_dim", 64), critic_hidden_dim=cfg.get("critic_hidden_dim", 64), @@ -81,6 +82,7 @@ def _make_td3_runner(env_name: str) -> OffPolicyRunner: learner = FastTD3Learner( obs_dim=obs_dim, action_dim=action_dim, + critic_obs_dim=obs_dim, num_envs=4, device="cpu", actor_hidden_dim=64, diff --git a/tests/utils/test_obs_utils.py b/tests/utils/test_obs_utils.py index 70d644eaf..1a2158c27 100644 --- a/tests/utils/test_obs_utils.py +++ b/tests/utils/test_obs_utils.py @@ -108,7 +108,8 @@ def test_no_critic(self): obs = {"obs": np.ones((4, 8))} obs_arr, critic_arr = split_obs_dict(obs) assert obs_arr.shape == (4, 8) - assert critic_arr is None + assert critic_arr.shape == (4, 8) + np.testing.assert_array_equal(critic_arr, obs_arr) class TestGetObsDims: @@ -124,4 +125,4 @@ def test_no_critic(self): spec = {"obs": 49} obs_dim, critic_dim = get_obs_dims(spec) assert obs_dim == 49 - assert critic_dim == 0 + assert critic_dim == 49 From 39aac9776a45e6a7a00721ebe1e7d829801b78fe Mon Sep 17 00:00:00 2001 From: tatp-yf Date: Wed, 22 Apr 2026 15:10:21 +0800 Subject: [PATCH 3/4] fix: remove unused variable flagged by ruff (F841) --- src/unilab/algos/torch/fast_td3/learner.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/unilab/algos/torch/fast_td3/learner.py b/src/unilab/algos/torch/fast_td3/learner.py index 3c4a7b268..0e6aab273 100644 --- a/src/unilab/algos/torch/fast_td3/learner.py +++ b/src/unilab/algos/torch/fast_td3/learner.py @@ -249,7 +249,7 @@ def normalize_obs(self, obs: torch.Tensor, update: bool = False) -> torch.Tensor def update_critic(self, data: Dict[str, torch.Tensor]) -> Dict[str, float]: """One critic update step.""" - observations = self.normalize_obs(data["obs"], update=True) + self.normalize_obs(data["obs"], update=True) critic_obs = data["critic"] actions = data["actions"] rewards = data["rewards"] From 3151b4f4af135cfc9d230c11fec201fdd90d2858 Mon Sep 17 00:00:00 2001 From: tatp-yf Date: Wed, 22 Apr 2026 15:53:28 +0800 Subject: [PATCH 4/4] style: apply ruff format --- src/unilab/algos/torch/fast_td3/runner.py | 4 +++- src/unilab/utils/final_observation.py | 1 + 2 files changed, 4 insertions(+), 1 deletion(-) diff --git a/src/unilab/algos/torch/fast_td3/runner.py b/src/unilab/algos/torch/fast_td3/runner.py index 77106d74e..8a6775467 100644 --- a/src/unilab/algos/torch/fast_td3/runner.py +++ b/src/unilab/algos/torch/fast_td3/runner.py @@ -42,7 +42,9 @@ def __init__( obs_normalization: bool = True, sim_backend: str = "mujoco", ): - obs_dim, action_dim, critic_obs_dim = get_env_dims(env_name, sim_backend, env_cfg_override=env_cfg_override) + obs_dim, action_dim, critic_obs_dim = get_env_dims( + env_name, sim_backend, env_cfg_override=env_cfg_override + ) learner = FastTD3Learner( obs_dim=obs_dim, action_dim=action_dim, diff --git a/src/unilab/utils/final_observation.py b/src/unilab/utils/final_observation.py index 08937f593..f18fc1ba6 100644 --- a/src/unilab/utils/final_observation.py +++ b/src/unilab/utils/final_observation.py @@ -5,6 +5,7 @@ import numpy as np + @dataclass(frozen=True) class TransitionBootstrapContract: actor_next_obs: np.ndarray